From bcfead8f0a534ab3bb04902fdb306edd494948ed Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8D=97=E9=9C=84?= Date: Wed, 14 Jan 2026 16:26:38 +0800 Subject: [PATCH 01/11] PullRequest: 5 core MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Merge branch core of git@code.alipay.com:pia/linghe.git into main https://code.alipay.com/pia/linghe/pull_requests/5 * add mxfp8 quant * add mxfp8 quant * add mxfp8 quant * add mxfp8 quant * add mxfp8 quant * add mxfp8 quant * add mxfp8 quant * revise for colwise * opt mxfp8 quant * compatible with triton 3.2.0 * add arg silu for rope * fix norm * add fused block permute & unpermute * add fused permute/unpermute mxfp8 quantize * refine testcase * add batch scale & batch clip kernels * finetune params * add mla_rope kernel * add rope inferface & use float32 update in gemm * add quantization interface * add varlen rope * add cu_seqlens_kv for mla rope * add is_contiguous assert for rope * refine gemm testcase * use async h2d * add rope with cp * add rope with cp * add rope with cp * refine testcases * add support for mini v3 * add cp rope in facade * remove megatron code in test_rope * use two kernels for block_rms_norm * use two kernels for block_rms_norm * refine save_for_backward * use raw input instead of view in saved_tensor * fix bench norm bug * remove ctx.rms in norm * use parameter instead of parameter.data in op * add silu arg in rope * opt inplace impl in ce loss * refine arg names * add more barrier in ce loss * add assert in ce loss * use meaningful shape for batch quantization * support transpose in mla * support transpose in mla * fix arg num bug in mla rope * linghe 0.1.0 * add approvers in aci * support more num_experts in fp32_gemm * update aci * update * refine testcases * remove is_contiguous assertion in batch norm * use fp32 dw in multiple kernels * use 0.1.2 version * recover test image * use int64 in clip and scale and fp32 in rope * add is_contiguous assertion and use rsqrt instead of 1/sqrt * mxfp8 deepep permute * mxfp8 unpermute func * support v3 dim * assert * scatter add for v3 funny hidden * use numerical stable impl for ce * use float64 in ce * refine testcases * refine testcases * refine testcases * use fp32 in silu and split batch silu backward * fix bug in topk * revise&refine test * tune params * add arg reuse in rope * adapt for triton 3.2.0 * support tp in rope * support tp in rope * use int64 in batch ops * use int64 in batch ops * use mask instead of max to accelerate clip * mv mxfp8 silu to linghe * mv mxfp8 rms to linghe * gradient fix * refine args in mxfp8 * use int64 in smooth/mxfp8 batch kernels * fix transpose in mla * fix transpose in mla * fix transpose in mla * fix transpose in mla * fix transpose in mla * fix transpose in mla * fix transpose in mla * fix mla bug with bs=1 and strided k_pos_emb * use padding batch instead of 128 in varlen rope * do not transpose mla rope output with thd layout * support more shapes * refine test.sh * remove rms none assertion in rms norm * add embedding and mla * refine mla impl * add for hybridep * add varlen mla * use 0.2.5 * use 0.2.5 * return zero when grad tensors is empty * define inf/inf=1 * add numel>0 assert in batch ops * add experiment op&reduce confusion * WIP dist ce loss * refine testcases * add debug log * add debug log * support ignore_index in ce * add parallel ce * add parallel ce * version 0.2.7 * remove inf_or_nan * add log for labels * add barrier in ce backward * refine testcase * add tp_group in ce * add tp_group in ce * use rsqrt instead of 1/sqrt * use 0.2.8 * refine condition of ce parallel * support stride for grad of embedding * use 0.2.9 * return grad for dummy tensor in embedding * support bf16 in batch ops * use fast impl for embedding * add ptx util & revise count zeros * fix typo in triton_batch_count_zero * use fast and accurate impl for embedding * refactor embedding ops * refactor embedding ops * revise for liuyu request * PullRequest: 2 合并计算通信融合算子 * PullRequest: 3 reformat * refine for public --- .aci.yml | 86 + CHANGELOG.md | 34 + benchmark/bench_gemm.py | 258 +- benchmark/bench_grad_norm.py | 41 + benchmark/bench_la.py | 35 + benchmark/bench_loss.py | 111 +- benchmark/bench_mla.py | 582 ++++ benchmark/bench_mla_rope.py | 272 ++ benchmark/bench_norm.py | 92 + benchmark/bench_permutation.py | 107 +- benchmark/bench_quantization.py | 107 + benchmark/bench_rmsnorm.py | 64 - benchmark/bench_topk.py | 108 + linghe/attn/__init__.py | 0 linghe/attn/la.py | 1590 +++++++++ linghe/attn/mla.py | 1968 ++++++++++++ linghe/experimental/__init__.py | 3 + linghe/experimental/demb.py | 536 +++ linghe/experimental/dla.py | 1197 +++++++ linghe/experimental/dmm.py | 195 ++ .../experimental/gmem_barrier_arrive_wait.py | 74 + linghe/experimental/norm.py | 352 ++ linghe/experimental/symm_mem_barrier.py | 165 + linghe/experimental/test_demb.py | 208 ++ linghe/experimental/test_dla.py | 177 + linghe/experimental/test_dmm.py | 101 + linghe/experimental/test_norm.py | 87 + linghe/facade/add.py | 3 +- linghe/facade/emb.py | 109 + linghe/facade/fp32_gemm.py | 31 +- linghe/facade/gate.py | 63 + linghe/facade/hadamard_quant_linear.py | 41 +- linghe/facade/loss.py | 94 +- linghe/facade/mla.py | 116 + linghe/facade/norm.py | 137 +- linghe/facade/permutation.py | 331 ++ linghe/facade/rope.py | 232 +- linghe/facade/silu.py | 160 + linghe/facade/smooth_quant_linear.py | 45 +- linghe/facade/topk.py | 99 + linghe/facade/transpose.py | 20 +- linghe/gemm/blockwise_fp8_gemm.py | 40 +- linghe/gemm/fp32_gemm.py | 326 +- linghe/quant/block.py | 226 +- linghe/quant/channel.py | 2 +- linghe/quant/group.py | 36 +- linghe/quant/hadamard.py | 6 +- linghe/quant/smooth.py | 177 +- linghe/tools/__init__.py | 0 linghe/tools/benchmark.py | 8 +- linghe/tools/check.py | 145 + linghe/tools/util.py | 100 +- linghe/utils/add.py | 4 +- linghe/utils/dot.py | 57 - linghe/utils/emb.py | 458 +++ linghe/utils/gate.py | 259 ++ linghe/utils/gather.py | 237 +- linghe/utils/loss.py | 481 ++- linghe/utils/mul.py | 158 + linghe/utils/norm.py | 807 +++-- linghe/utils/rearange.py | 15 +- linghe/utils/reduce.py | 182 +- linghe/utils/rope.py | 2859 +++++++++++++++-- linghe/utils/scatter.py | 62 +- linghe/utils/silu.py | 1027 ++++-- linghe/utils/topk.py | 289 ++ linghe/utils/transpose.py | 109 +- linghe/utils/unary.py | 82 +- scripts/dev.py | 30 + scripts/plot_input_output.py | 74 +- scripts/reproduce_triton_bug.py | 61 +- scripts/test.sh | 44 +- setup.py | 2 +- tests/test_add.py | 5 +- tests/test_blockwise_fp8_gemm.py | 51 +- tests/test_blockwise_quant.py | 115 + tests/test_channel_quant.py | 16 +- tests/test_channelwise_fp8_gemm.py | 55 +- tests/test_dist_loss.py | 159 + tests/test_dot.py | 42 - tests/test_embedding.py | 155 + tests/test_fp32_gemm.py | 166 +- tests/test_gate.py | 140 + tests/test_gather.py | 482 ++- tests/test_group_quant.py | 12 +- tests/test_hadamard_quant.py | 85 +- tests/test_la.py | 200 ++ tests/test_loss.py | 177 +- tests/test_mla.py | 352 ++ tests/test_mul.py | 104 + tests/test_mxfp8_quant.py | 98 + tests/test_norm.py | 327 +- tests/test_rearange.py | 40 +- tests/test_reduce.py | 79 +- tests/test_rope.py | 641 +++- tests/test_scatter.py | 25 +- tests/test_silu.py | 515 ++- tests/test_smooth_quant.py | 109 +- tests/test_topk.py | 211 ++ tests/test_transpose.py | 60 +- tests/test_unary.py | 81 +- 101 files changed, 20188 insertions(+), 2708 deletions(-) create mode 100644 .aci.yml create mode 100644 CHANGELOG.md create mode 100644 benchmark/bench_grad_norm.py create mode 100644 benchmark/bench_la.py create mode 100644 benchmark/bench_mla.py create mode 100644 benchmark/bench_mla_rope.py create mode 100644 benchmark/bench_norm.py create mode 100644 benchmark/bench_quantization.py delete mode 100644 benchmark/bench_rmsnorm.py create mode 100644 benchmark/bench_topk.py create mode 100644 linghe/attn/__init__.py create mode 100644 linghe/attn/la.py create mode 100644 linghe/attn/mla.py create mode 100644 linghe/experimental/__init__.py create mode 100644 linghe/experimental/demb.py create mode 100644 linghe/experimental/dla.py create mode 100644 linghe/experimental/dmm.py create mode 100644 linghe/experimental/gmem_barrier_arrive_wait.py create mode 100644 linghe/experimental/norm.py create mode 100644 linghe/experimental/symm_mem_barrier.py create mode 100644 linghe/experimental/test_demb.py create mode 100644 linghe/experimental/test_dla.py create mode 100644 linghe/experimental/test_dmm.py create mode 100644 linghe/experimental/test_norm.py create mode 100644 linghe/facade/emb.py create mode 100644 linghe/facade/gate.py create mode 100644 linghe/facade/mla.py create mode 100644 linghe/facade/permutation.py create mode 100644 linghe/facade/silu.py create mode 100644 linghe/facade/topk.py create mode 100644 linghe/tools/__init__.py create mode 100644 linghe/tools/check.py delete mode 100644 linghe/utils/dot.py create mode 100644 linghe/utils/emb.py create mode 100644 linghe/utils/gate.py create mode 100644 linghe/utils/mul.py create mode 100644 linghe/utils/topk.py create mode 100644 scripts/dev.py create mode 100644 tests/test_blockwise_quant.py create mode 100644 tests/test_dist_loss.py delete mode 100644 tests/test_dot.py create mode 100644 tests/test_embedding.py create mode 100644 tests/test_gate.py create mode 100644 tests/test_la.py create mode 100644 tests/test_mla.py create mode 100644 tests/test_mul.py create mode 100644 tests/test_mxfp8_quant.py create mode 100644 tests/test_topk.py diff --git a/.aci.yml b/.aci.yml new file mode 100644 index 0000000..f3dceb3 --- /dev/null +++ b/.aci.yml @@ -0,0 +1,86 @@ +version: "2.0" + +stages: +- 前置检查 +- 构建&发布到测试库 +- 验包&确认 +- 发布正式库前检查 +- 发布正式库 + +jobs: + 单元测试: + stage: 前置检查 + component: python-ut + inputs: + languageConfig: + pythonVersion: "3.9.0" + config: + execute: + isAllowSkip: true + + 代码检查: + stage: 前置检查 + component: python-sast + inputs: + excludes: # 选填项,排除哪些项不进行代码扫描 + - "**__init__.py**" + - "**/tests/**" + - ansible/* + - config/* + - benchmark/* + - scripts/* + - examples/* + - docs/* + codePath: "./" # 选填项,选择扫描目录 + config: + execute: + timeout: 600 # 选填项,任务超时时间 + isAllowSkip: true + afterExecute: + checkRule: # 选填项,卡点策略, 根据实际团队质量要求进行配置 + - ${{outputs.critical}} <= 500 + + STC安全扫描: + stage: 前置检查 + component: stc + inputs: + tenantName: null + config: + execute: + isAllowSkip: true + + 构建并发布到测试库: + stage: 构建&发布到测试库 + id: build + component: pypi-artifact-uploader + inputs: + buildImage: reg.docker.alibaba-inc.com/aii/aistudio:aistudio-190677225-3221750112-1752554942251 + # buildImage: reg.docker.alibaba-inc.com/aii/aistudio:12150173-20251107143737 # max v2 + # buildTool: poetry + # artifactType: "wheel" # 仅打wheel包,如需要 tgz,请删除此行 + registry: "https://artifacts.antgroup-inc.cn/artifact/repositories/simple-dev/" # 测试库地址 + # workdir: . # pypi 工程目录,默认在此目录下进行 python -m build 并输出到 dist 目录, 详细可参考组件首页说明 + buildCmd: "python setup.py bdist_wheel" + only: + - master + + 选择迁移制品: + id: check + stage: 发布正式库前检查 + component: artifact-transfer-check + inputs: + artifactsConfigs: + - artifacts: ${{jobs.build.outputs.artifacts}} + antArtifactRepo: simple + + 同步至正式库: + stage: 发布正式库 + component: ant-artifact-transfer + inputs: + transferArtifacts: ${{jobs.check.outputs.transferArtifacts}} + config: + beforeExecute: + isAutoSkip: ${{jobs.check.outputs.selectedCount}} = 0 + confirm: + buttonName: 确认发布正式库 + approvers: ["nanxiao.zy", "liangchen.liangche"] # 具备发布角色的用户的域账号, 如 ["yunjie.gyj"] 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/benchmark/bench_gemm.py b/benchmark/bench_gemm.py index 862e7d9..17ff702 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, @@ -38,8 +36,7 @@ def torch_accum_weight(x, w, out, x_scale, w_scale): 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' n_repeat = 100 @@ -58,73 +55,101 @@ def test_cublas_blockwise_gemm(M=4096, N=4096, K=4096): 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}') + 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}') + 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}') + 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 +167,114 @@ 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) + 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 + 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) - + # 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..d3611e0 --- /dev/null +++ b/benchmark/bench_grad_norm.py @@ -0,0 +1,41 @@ +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..c2cb67a --- /dev/null +++ b/benchmark/bench_la.py @@ -0,0 +1,35 @@ +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..e6047cc 100644 --- a/benchmark/bench_loss.py +++ b/benchmark/bench_loss.py @@ -11,98 +11,97 @@ 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.tools.check import output_check from linghe.utils.loss import (triton_softmax_cross_entropy_backward, - triton_softmax_cross_entropy_forward) + 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() +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 te_cross_entropy_forward_backward(logits, targets): +def te_cross_entropy_forward_backward(logits, targets, input_grad): losses = parallel_cross_entropy(logits[None], targets[None]) - loss = losses.mean() - loss.backward() + losses.backward(input_grad[None]) 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 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 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 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 + 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) 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) + 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, - ref_bytes=M * N * 4, + input_grad, + n_repeat=1, + ref_bytes=M * N * 6, ref_time=ref_time) - benchmark_func(te_cross_entropy_forward_backward, - logits.detach().clone().requires_grad_(), targets, - ref_bytes=M * N * 4, + 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_cross_entropy_forward_backward, logits, targets, - input_grad, 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) 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)) + 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..3063697 --- /dev/null +++ b/benchmark/bench_mla.py @@ -0,0 +1,582 @@ +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..a2aecf4 --- /dev/null +++ b/benchmark/bench_mla_rope.py @@ -0,0 +1,272 @@ +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..45042cb --- /dev/null +++ b/benchmark/bench_norm.py @@ -0,0 +1,92 @@ +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..7b86295 100644 --- a/benchmark/bench_permutation.py +++ b/benchmark/bench_permutation.py @@ -1,15 +1,19 @@ 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,11 +33,52 @@ 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' 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) @@ -50,7 +95,7 @@ def bench_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8): 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) + 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, @@ -58,6 +103,32 @@ def bench_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8): 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) + + 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) + + 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' @@ -73,7 +144,7 @@ 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 @@ -81,16 +152,18 @@ def bench_triton_unpermute_with_mask_map(M=4098, N=4096, n_experts=32, topk=2): 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) + 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) + 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) +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..3515963 --- /dev/null +++ b/benchmark/bench_quantization.py @@ -0,0 +1,107 @@ +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..377c7cd --- /dev/null +++ b/benchmark/bench_topk.py @@ -0,0 +1,108 @@ +# -*- 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/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..0d6af35 --- /dev/null +++ b/linghe/attn/la.py @@ -0,0 +1,1590 @@ +# -*- 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..6534f02 --- /dev/null +++ b/linghe/attn/mla.py @@ -0,0 +1,1968 @@ +# -*- 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 + 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 + 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 + 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) + 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 + 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 = max(triton.next_power_of_2(B), 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..b86edf4 --- /dev/null +++ b/linghe/experimental/demb.py @@ -0,0 +1,536 @@ +# -*- 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..cd45da7 --- /dev/null +++ b/linghe/experimental/dla.py @@ -0,0 +1,1197 @@ +# -*- 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..20381f8 --- /dev/null +++ b/linghe/experimental/dmm.py @@ -0,0 +1,195 @@ +# -*- 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/gmem_barrier_arrive_wait.py b/linghe/experimental/gmem_barrier_arrive_wait.py new file mode 100644 index 0000000..b98d983 --- /dev/null +++ b/linghe/experimental/gmem_barrier_arrive_wait.py @@ -0,0 +1,74 @@ +""" +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/norm.py b/linghe/experimental/norm.py new file mode 100644 index 0000000..d69b66b --- /dev/null +++ b/linghe/experimental/norm.py @@ -0,0 +1,352 @@ +# -*- 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 + +""" +the code is used to reproduce the barrier bug in cross-block reduce. +""" + + +@triton.jit +def rms_norm_forward_kernel(x_ptr, + weight_ptr, + out_ptr, + cache_ptr, + signal_ptr, + rms_ptr, + eps, + M, + n: tl.constexpr, + H: tl.constexpr, + B: tl.constexpr): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + CB = tl.num_programs(1) + + indices = rid * H + tl.arange(0, H) + offs = rid * H * n + cid * B + tl.arange(0, H)[:, None] * n + tl.arange(0, + B)[ + None, :] + + x = tl.load(x_ptr + offs).to(tl.float32) + weight = tl.load(weight_ptr + cid * B + tl.arange(0, B)).to(tl.float32) + + s = tl.sum(x * x, axis=1) + tl.atomic_add(cache_ptr + indices, s, sem='acq_rel', scope='sys') + # tl.debug_barrier() + tl.atomic_add(signal_ptr + rid, 1, sem='acq_rel', scope='sys') + tl.debug_barrier() + # tl.inline_asm_elementwise( + # "membar.gl;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 + # ) + # tl.inline_asm_elementwise( + # "bar.sync 0;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 + # ) + count = tl.load(signal_ptr + rid, cache_modifier='.cv') + while count < CB: + count = tl.load(signal_ptr + rid, cache_modifier='.cv') + # if cid + rid == 0: + # tl.device_print('count', count) + tl.debug_barrier() + + sums = tl.load(cache_ptr + indices, cache_modifier='.cv', volatile=True) + # sums = tl.atomic_add(cache_ptr + indices, 0.0, sem='acq_rel', scope='gpu') + + rms = tl.rsqrt(sums / n + eps) + if cid == 0: + tl.store(rms_ptr + indices, rms) + + x = x * rms[:, None] * weight[None, :] + + tl.store(out_ptr + offs, x) + + +def triton_rms_norm_forward(x: torch.Tensor, + weight: torch.Tensor, + eps: float = 1e-6): + """ + Fused RMSNorm forward. + Args: + x: Input tensor, shape [M, N] + weight: RMSNorm weight, shape [N] + eps: epsilon value for L2 normalization. + Returns: + - out: output. + - rms: Reciprocal of the root mean square of the + input calculated over the last dimension. + """ + assert x.is_contiguous() and weight.is_contiguous() + M, n = x.shape + device = x.device + + out = torch.empty((M, n), dtype=x.dtype, device=device) + rms = torch.empty((M,), dtype=torch.float32, device=device) + H = 8 + B = 128 + assert M % H == 0 and n % B == 0 + CB = n // 128 # column block + RB = M // H # row block + cache = torch.zeros((M,), dtype=torch.float32, device=device) + signals = torch.zeros((RB,), dtype=torch.int32, device=device) + grid = (RB, CB) + rms_norm_forward_kernel[grid]( + x, + weight, + out, + cache, + signals, + rms, + eps, + M, + n, + H, + B, + num_stages=1, + num_warps=1 + ) + return out, rms + + +@triton.jit +def _parallel_rms_norm_and_block_quant_forward_kernel(x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + cache_ptr, + rms_ptr, + eps, + M, + n: tl.constexpr, + H: tl.constexpr, + B: tl.constexpr, + K: tl.constexpr, + ROUND: tl.constexpr): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + CB = tl.num_programs(1) + + indices = rid * H + tl.arange(0, H) + masks = indices[:, None] < M + offs = rid * H * n + cid * B + tl.arange(0, H)[:, None] * n + tl.arange(0, + B)[ + None, :] + + x = tl.load(x_ptr + offs, mask=masks).to(tl.float32) + s = tl.sum(x * x, axis=1) + # tl.debug_barrier() + tl.atomic_add(cache_ptr + indices, s, scope='sys') + # tl.debug_barrier() + tl.atomic_add(cache_ptr + M + rid, 1.0, scope='sys') + # tl.inline_asm_elementwise( + # "bar.sync 0;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 + # ) + tl.debug_barrier() + # count = tl.atomic_add(cache_ptr + M + rid, 0.0) + # tl.inline_asm_elementwise( + # "membar.gl;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 + # ) + for i in range(3): + count = tl.load(cache_ptr + M + rid) + # tl.debug_barrier() + while count < CB: + count = tl.load(cache_ptr + M + rid) + tl.debug_barrier() + # tl.inline_asm_elementwise( + # "membar.gl;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 + # ) + sums = tl.load(cache_ptr + indices, mask=indices < M) + + rms = tl.rsqrt(sums / n + eps) + if cid == CB - 1: + tl.store(rms_ptr + indices, rms, mask=indices < M) + + toffs = cid * M * B + rid * H + tl.arange(0, B)[:, None] * M + tl.arange(0, + H)[ + None, :] + weight = tl.load(weight_ptr + cid * B + tl.arange(0, B)).to(tl.float32) + x = x * rms[:, None] * weight[None, :] + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + q = (x / scale[:, None]).to(out_ptr.dtype.element_ty) + + tl.store(scale_ptr + cid * M + indices, scale, mask=indices < M) + tl.store(out_ptr + offs, q, mask=masks) + + 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 * B + tl.arange(0, B), scale) + + q = (tl.trans(x / scale)).to(transpose_output_ptr.dtype.element_ty) + tl.store(transpose_output_ptr + toffs, q, mask=indices[None, :] < M) + + +@triton.jit +def parallel_rms_norm_and_block_quant_forward_kernel(x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + cache_ptr, + rms_ptr, + eps, + M, + n: tl.constexpr, + H: tl.constexpr, + B: tl.constexpr, + K: tl.constexpr, + ROUND: tl.constexpr): + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + CB = tl.num_programs(1) + + indices = rid * H + tl.arange(0, H) + masks = indices[:, None] < M + offs = rid * H * n + cid * K * B + tl.arange(0, H)[:, None] * n + tl.arange( + 0, B)[ + None, :] + + s = tl.zeros((H,), dtype=tl.float32) + for i in range(K): + x = tl.load(x_ptr + i * B + offs, mask=masks).to(tl.float32) + s += tl.sum(x * x, axis=1) + + tl.atomic_add(cache_ptr + indices, s) + tl.atomic_add(cache_ptr + M + rid, 1.0) + tl.debug_barrier() + + count = tl.atomic_add(cache_ptr + M + rid, 0.0) + # count = tl.load(cache_ptr + M + rid) + tl.debug_barrier() + while count < 1.0 * CB: + count = tl.load(cache_ptr + M + rid) + for i in range(1024): + pass + tl.debug_barrier() + sums = tl.load(cache_ptr + indices, mask=indices < M) + + rms = tl.rsqrt(sums / n + eps) + + if cid == CB - 1: + tl.store(rms_ptr + indices, rms, mask=indices < M) + + toffs = cid * M * B * K + rid * H + tl.arange(0, B)[:, + None] * M + tl.arange(0, H)[ + None, :] + for i in range(K): + weight = tl.load(weight_ptr + cid * K * B + i * B + tl.arange(0, B)).to( + tl.float32) + x = tl.load(x_ptr + i * B + offs, mask=masks).to(tl.float32) + x = x * rms[:, None] * weight[None, :] + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) + q = (x / scale[:, None]).to(out_ptr.dtype.element_ty) + + tl.store(scale_ptr + cid * K * M + i * M + indices, scale, + mask=indices < M) + tl.store(out_ptr + i * B + offs, q, mask=masks) + + 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 * B * K + i * B + tl.arange(0, + B), + scale) + + q = (tl.trans(x / scale)).to(transpose_output_ptr.dtype.element_ty) + tl.store(transpose_output_ptr + i * B * M + toffs, q, + mask=indices[None, :] < M) + + +def triton_parallel_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: + x: Input tensor, shape [M, N] + weight: RMSNorm weight, shape [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. + output_mode: one of {0, 1, 2}. + 0: only output non-transpose tensor + 1: only output transposed tensor + 2: return both + Returns: + - out: quantization data. + - scale: quantization scale. + - rms: Reciprocal of the root mean square of the + input calculated over the last dimension. + - transpose_output: quantization data of transposed gradient. + - transpose_scale: quantization scale of transposed gradient. + """ + # row-wise read, row-wise write + assert x.is_contiguous() and weight.is_contiguous() + M, n = x.shape + 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((n // 128, M), device=device, dtype=torch.float32) + + 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) + + assert rms is None + rms = torch.empty((M,), dtype=torch.float32, device=device) + H = 128 + B = 128 + CB = n // B + K = n // (B * CB) + assert K >= 1 + RB = triton.cdiv(M, H) + cache = torch.zeros((M + RB,), dtype=torch.float32, device=device) + grid = (RB, CB) + _parallel_rms_norm_and_block_quant_forward_kernel[grid]( + x, + weight, + out, + scale, + transpose_output, + transpose_scale, + cache, + rms, + eps, + M, + n, + H, + B, + K, + round_scale, + num_stages=5, + num_warps=4 + ) + return out, scale, rms, transpose_output, transpose_scale diff --git a/linghe/experimental/symm_mem_barrier.py b/linghe/experimental/symm_mem_barrier.py new file mode 100644 index 0000000..666df67 --- /dev/null +++ b/linghe/experimental/symm_mem_barrier.py @@ -0,0 +1,165 @@ +# -*- 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..b26c241 --- /dev/null +++ b/linghe/experimental/test_demb.py @@ -0,0 +1,208 @@ +# -*- 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..d67f8b1 --- /dev/null +++ b/linghe/experimental/test_dla.py @@ -0,0 +1,177 @@ +# -*- 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..07d90bf --- /dev/null +++ b/linghe/experimental/test_dmm.py @@ -0,0 +1,101 @@ +# -*- 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/experimental/test_norm.py b/linghe/experimental/test_norm.py new file mode 100644 index 0000000..95074f8 --- /dev/null +++ b/linghe/experimental/test_norm.py @@ -0,0 +1,87 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.experimental.norm import triton_rms_norm_forward, \ + triton_parallel_rms_norm_and_block_quant_forward +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check +from linghe.utils.norm import triton_rms_norm_and_block_quant_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 + ) + with torch.no_grad(): + rmsnorm.weight.copy_(weight) + rms = torch.rsqrt(torch.sum(x ** 2, 1) / N + 1e-6) + return rmsnorm(x).to(dtype), rms + + +def test_norm(M=4096, N=4096, bench=False): + dtype = torch.bfloat16 + device = 'cuda:0' + + x = torch.ones(M, N, dtype=dtype, requires_grad=False, device=device) + weight = torch.ones(N, dtype=dtype, requires_grad=False, device=device) + + output_ref, rms_ref = torch_rms_forward(x, weight) + output, rms = triton_rms_norm_forward(x, weight) + + output_check(rms_ref, rms, name="rms", rtol=0.001) + output_check(output_ref, output, name="output", rtol=0.001) + + +def test_parallel_rmsnorm_and_block_quant(M=4096, N=4096, bench=False): + dtype = torch.bfloat16 + 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) + + q_ref, scale_ref, rms_ref, qt_ref, scale_t_ref = triton_rms_norm_and_block_quant_forward( + x, weight, + round_scale=False, + output_mode=2) + + q, scale, rms, qt, scale_t = triton_parallel_rms_norm_and_block_quant_forward( + x, weight, + round_scale=False, + output_mode=2) + output_check(q_ref, q, name='parallel.block.data', rtol=-0.125) + output_check(scale_ref, scale, name="parallel.block.scale", rtol=-0.125) + output_check(rms_ref, rms, name="parallel.block.rms", rtol=-0.125) + output_check(qt_ref, qt, name='parallel.block.t_data', rtol=-0.125) + output_check(scale_t_ref, scale_t, name="parallel.block.t_scale", + rtol=-0.125) + + if bench: + benchmark_func(triton_rms_norm_and_block_quant_forward, x, weight, + round_scale=False, + output_mode=2, + ref_bytes=M * N * 4, + n_profile=2) + + benchmark_func(triton_parallel_rms_norm_and_block_quant_forward, x, + weight, + round_scale=False, + output_mode=2, + ref_bytes=M * N * 4, + n_profile=2) + + +if __name__ == '__main__': + # /usr/local/lib/python3.12/dist-packages/triton/backends/nvidia/bin/ptxas -lineinfo -v --gpu-name=sm_90a /tmp/tmp3l_m5rfp.ptx -o /tmp/tmp3l_m5rfp.ptx.o + test_norm(M=1024, N=4096, bench=False) + # test_parallel_rmsnorm_and_block_quant(M=4096, N=4096, bench=False) 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..e0bdaed --- /dev/null +++ b/linghe/facade/emb.py @@ -0,0 +1,109 @@ +# -*- 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..20c79cf 100644 --- a/linghe/facade/fp32_gemm.py +++ b/linghe/facade/fp32_gemm.py @@ -6,35 +6,40 @@ import torch from linghe.gemm.fp32_gemm import (triton_fp32_gemm, - triton_fp32_gemm_for_backward, - triton_fp32_gemm_for_update) + 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) + if len(shape) == 3: + input = input.view(shape[0] * shape[1], shape[2]) + logits = triton_fp32_gemm(input, weight) 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]) + 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): - shape = grad_output.shape - grad_output = grad_output.view(shape[0] * shape[1], shape[2]) + 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 dx = triton_fp32_gemm_for_backward(grad_output, weight) - dx = dx.view(*ctx.shape) + if len(grad_shape) == 3: + dx = dx.view(*ctx.shape) dw = triton_fp32_gemm_for_update(grad_output, input) @@ -51,4 +56,4 @@ def fp32_gemm(input: torch.Tensor, weight: torch.Tensor): Returns: output of gemm """ - return Fp32GEMM.apply(input, weight) \ No newline at end of file + return Fp32GEMM.apply(input, weight) diff --git a/linghe/facade/gate.py b/linghe/facade/gate.py new file mode 100644 index 0000000..1d5e02f --- /dev/null +++ b/linghe/facade/gate.py @@ -0,0 +1,63 @@ +# -*- 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 + + +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, + ctx.eps, + 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): + """ + 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) diff --git a/linghe/facade/hadamard_quant_linear.py b/linghe/facade/hadamard_quant_linear.py index 5b3dd45..e017ea6 100644 --- a/linghe/facade/hadamard_quant_linear.py +++ b/linghe/facade/hadamard_quant_linear.py @@ -11,7 +11,6 @@ from linghe.quant.hadamard import triton_hadamard_quant - class _HadamardQuantLinear(torch.autograd.Function): @staticmethod def forward( @@ -29,13 +28,15 @@ def forward( 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) + 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), + scale_a=x_scale.view(-1, 1), + scale_b=w_scale.view(1, -1), out_dtype=ctx.out_dtype, use_fast_accum=True ) @@ -64,25 +65,26 @@ def backward( 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 - ) + 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 - ) + 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,6 +97,7 @@ class HadamardQuantLinear(torch.nn.Module): """ a naive implementation of hadamard transformation and quantization """ + def __init__( self, in_features: int, @@ -146,7 +149,7 @@ def forward(self, input: torch.Tensor) -> torch.Tensor: """""" if self.training: return _HadamardQuantLinear.apply(input, self.weight, self.bias, - self.hadamard_matrix) + self.hadamard_matrix) else: output = input @ self.weight.t() if self.bias is not None: diff --git a/linghe/facade/loss.py b/linghe/facade/loss.py index 0feac8a..d1d1abc 100644 --- a/linghe/facade/loss.py +++ b/linghe/facade/loss.py @@ -5,22 +5,37 @@ 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, 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.parallel = parallel if len(shape) == 3: loss = loss.view(shape[0], shape[1]) return loss @@ -29,17 +44,34 @@ 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,11 +83,13 @@ 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): """""" + @staticmethod def forward(ctx, x, coef=0.2): ctx.coef = coef @@ -70,3 +104,33 @@ def backward(ctx, grad_output): scale = 1 / torch.pow(array.float(), ctx.coef) grad = grad_output * scale return grad, None + + +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..59b7640 --- /dev/null +++ b/linghe/facade/mla.py @@ -0,0 +1,116 @@ +# -*- 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..be90bf7 100644 --- a/linghe/facade/norm.py +++ b/linghe/facade/norm.py @@ -5,36 +5,27 @@ 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, +) 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 +40,76 @@ 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): - """""" - @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 - ) - ctx.save_for_backward(attn_output, gate, weight.data) - ctx.eps = eps - ctx.group_size = group_size - - return output +# used in attention rms norm +class BlockRMSNorm(torch.autograd.Function): @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 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, + ) ) - 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( + 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.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 + 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 diff --git a/linghe/facade/permutation.py b/linghe/facade/permutation.py new file mode 100644 index 0000000..d044a2e --- /dev/null +++ b/linghe/facade/permutation.py @@ -0,0 +1,331 @@ +from typing import Optional, List + +import torch + +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, +) +from linghe.utils.scatter import triton_unpermute_with_mask_map + + +class _PaddedPermute(torch.autograd.Function): + @staticmethod + def forward( + ctx, + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + ): + """Forward function.""" + num_tokens, hidden_dim = tokens.shape + + row_id_map = triton_make_row_id_map(routing_map, multiple_of=16) + num_out_tokens = sum( + [(x + 15) // 16 * 16 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, + ) + + +def padded_permute( + tokens, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + 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 = _PaddedPermute.apply( + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + ) + 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 diff --git a/linghe/facade/rope.py b/linghe/facade/rope.py index c31f1c8..0844496 100644 --- a/linghe/facade/rope.py +++ b/linghe/facade/rope.py @@ -3,80 +3,242 @@ 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 + 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): + 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] - - ko: shape [B, S, h, head_dim] - - vo: shape [B, S, h, head_dim] + - 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) \ No newline at end of file + 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 [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] + """ + 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..02f0317 --- /dev/null +++ b/linghe/facade/silu.py @@ -0,0 +1,160 @@ +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, + ) + + +class BlockSiluFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, input, quantizer, grad_quantizer, cls): + 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.save_for_backward(input) + + x_q, x_scale, xt_q, xt_scale = triton_silu_and_block_quant_forward( + input_view, + round_scale=quantizer.force_pow_2_scales) + 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) + x_q, x_scale, xt_q, xt_scale = triton_silu_and_block_quant_backward( + grad_output_view, + input_view, + round_scale=grad_quantizer.force_pow_2_scales) + 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 + + +def block_silu_impl(input, quantizer, grad_quantizer, cls): + output = BlockSiluFunction.apply(input, quantizer, grad_quantizer, cls) + return output + + +class BlockBatchWeightedSiluFunction(torch.autograd.Function): + @staticmethod + def forward(ctx, input, weights, counts, splits, quantizers, + grad_quantizers, cls, 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.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_block_quant_forward(input, + weights, + counts, + splits=splits, + round_scale= + quantizers[ + 0].force_pow_2_scales, + 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 + (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=grad_quantizers[0].force_pow_2_scales) + 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 + + +def block_batch_weighted_silu_impl(input, weights, counts, splits, quantizers, + grad_quantizers, cls, is_recomputing=None): + assert input.ndim == 2 + output = BlockBatchWeightedSiluFunction.apply(input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + is_recomputing) + return output diff --git a/linghe/facade/smooth_quant_linear.py b/linghe/facade/smooth_quant_linear.py index dd284e2..52cf3de 100644 --- a/linghe/facade/smooth_quant_linear.py +++ b/linghe/facade/smooth_quant_linear.py @@ -7,11 +7,11 @@ 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.utils.reduce import triton_abs_max +from linghe.utils.transpose import triton_transpose_and_pad + class _SmoothQuantLinear(torch.autograd.Function): @staticmethod @@ -33,8 +33,10 @@ 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) + 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(), @@ -69,30 +71,30 @@ def backward( 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, + 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 , + 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) + 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,6 +107,7 @@ class SmoothQuantLinear(torch.nn.Module): """ a naive implementation of smooth quantization linear """ + def __init__( self, in_features: int, @@ -145,7 +148,7 @@ 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, diff --git a/linghe/facade/topk.py b/linghe/facade/topk.py new file mode 100644 index 0000000..40f4809 --- /dev/null +++ b/linghe/facade/topk.py @@ -0,0 +1,99 @@ +# -*- 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): + values, indices = triton_topk_forward(x, k, dim=dim) + 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 + + +def fused_topk(x, k, dim=-1): + """ + 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) + + +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..9e6ef5f 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..6de7cd2 100644 --- a/linghe/gemm/blockwise_fp8_gemm.py +++ b/linghe/gemm/blockwise_fp8_gemm.py @@ -3,19 +3,15 @@ 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, @@ -79,26 +75,25 @@ def triton_bb_fp8_gemm(a: torch.Tensor, 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 - ) + 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] +# ] -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.autotune(configs=fp8_gemm_configs, key=["N", "K"]) @triton.jit def fp8_gemm_tt_kernel( a_ptr, @@ -147,7 +142,6 @@ def fp8_gemm_tt_kernel( 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, @@ -165,5 +159,7 @@ def triton_tt_fp8_gemm(a: torch.Tensor, fp8_gemm_tt_kernel[grid](a, b, c, a_s, b_s, M, N, K, - BLOCK_SIZE_K=block_size) + BLOCK_SIZE_K=block_size, + BLOCK_SIZE_M=64, + BLOCK_SIZE_N=64) return c diff --git a/linghe/gemm/fp32_gemm.py b/linghe/gemm/fp32_gemm.py index 5ad778c..4f213b4 100644 --- a/linghe/gemm/fp32_gemm.py +++ b/linghe/gemm/fp32_gemm.py @@ -3,8 +3,6 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ -from typing import Optional - import torch import triton import triton.language as tl @@ -47,9 +45,8 @@ def fp32_gemm_kernel( 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 +56,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,19 +68,19 @@ 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) + assert x.is_contiguous() and w.is_contiguous() + M, K = x.size() + N, K = w.size() + assert M % 32 == 0 and K % 128 == 0 and N % 16 == 0 + 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 = 128 BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 + BLOCK_SIZE_N = max([x for x in [16, 32, 64, 128] if N % x == 0]) num_warps = 4 num_stages = 3 - fp32_gemm_kernel[grid](a, b, c, + fp32_gemm_kernel[grid](x, w, c, M, N, K, BLOCK_SIZE_K, BLOCK_SIZE_M, @@ -122,7 +118,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 +127,8 @@ 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,18 +137,18 @@ 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) + 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 = 128 + BLOCK_SIZE_K = max([x for x in [16, 32, 64, 128] if K % x == 0]) BLOCK_SIZE_M = 32 BLOCK_SIZE_N = 128 num_warps = 4 num_stages = 2 - fp32_gemm_for_backward_kernel[grid](a, b, c, + fp32_gemm_for_backward_kernel[grid](y, w, c, M, N, K, BLOCK_SIZE_K, BLOCK_SIZE_M, @@ -187,11 +182,9 @@ def fp32_gemm_for_update_kernel( 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,27 +195,27 @@ 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) + 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 = 128 - BLOCK_SIZE_M = 32 + BLOCK_SIZE_M = max([x for x in [16, 32] if M % x == 0]) BLOCK_SIZE_N = 128 num_warps = 4 num_stages = 3 - fp32_gemm_for_update_kernel[grid](a, b, c, + fp32_gemm_for_update_kernel[grid](y, x, c, M, N, K, BLOCK_SIZE_K, BLOCK_SIZE_M, @@ -233,12 +226,10 @@ def triton_fp32_gemm_for_update(a: torch.Tensor, b: torch.Tensor): return c - @triton.jit -def scaled_fp32_gemm_kernel( +def split_fp32_gemm_kernel( a_ptr, b_ptr, - scale_ptr, c_ptr, M, N: tl.constexpr, @@ -246,89 +237,172 @@ def scaled_fp32_gemm_kernel( BLOCK_SIZE_K: tl.constexpr, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, + SPLIT_COUNT: tl.constexpr ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) - k = tl.cdiv(K, BLOCK_SIZE_K) + pid_k = 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 + assert x.is_contiguous() and w.is_contiguous() + M, K = x.size() + N, K = w.size() BLOCK_SIZE_K = 128 - BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 + BLOCK_SIZE_M = 128 + BLOCK_SIZE_N = max([x for x in [16, 32, 64, 128] if N % x == 0]) + SPLIT_COUNT = 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 + + 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: (triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + SPLIT_COUNT) # noqa + 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 - ) + split_fp32_gemm_kernel[grid](x, w, c, + M, N, K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + SPLIT_COUNT, + num_warps=num_warps, + num_stages=num_stages + ) return c +# @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) +@triton.jit +def split_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, + SPLIT_COUNT: tl.constexpr +): + pid_m = tl.program_id(axis=0) + pid_n = tl.program_id(axis=1) + pid_k = 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() + 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 + SPLIT_COUNT = min(triton.cdiv(K, 2048), 8) + 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) + 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 = 2 + split_fp32_gemm_for_backward_kernel[grid](y, w, c, + M, N, K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + SPLIT_COUNT, + num_warps=num_warps, + num_stages=num_stages + ) + if SPLIT_COUNT > 1: + c = c.to(w.dtype) + return c + +# @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit -def scaled_fp32_gemm_for_update_kernel( +def split_fp32_gemm_for_update_kernel( a_ptr, b_ptr, - scale_ptr, c_ptr, M, N: tl.constexpr, @@ -336,23 +410,25 @@ def scaled_fp32_gemm_for_update_kernel( BLOCK_SIZE_K: tl.constexpr, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, + SPLIT_COUNT: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) - k = tl.cdiv(K, BLOCK_SIZE_K) + pid_k = 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 +436,45 @@ 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 + assert y.is_contiguous() and x.is_contiguous() + K, M = y.size() + K, N = x.size() + BLOCK_SIZE_K = 64 + BLOCK_SIZE_M = max([x for x in [16, 32, 64] if M % x == 0]) BLOCK_SIZE_N = 128 - num_warps = 4 + SPLIT_COUNT = min(triton.cdiv(K, 2048), 8) + 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: (triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + SPLIT_COUNT) # noqa + + num_warps = 2 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 - ) + split_fp32_gemm_for_update_kernel[grid](y, x, c, + M, N, K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + SPLIT_COUNT, + num_warps=num_warps, + num_stages=num_stages + ) return c diff --git a/linghe/quant/block.py b/linghe/quant/block.py index e7d093f..66cee97 100644 --- a/linghe/quant/block.py +++ b/linghe/quant/block.py @@ -29,10 +29,10 @@ def block_quant_kernel(x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, def triton_block_quant(x, - block_size=128, - round_scale=False): + 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,10 +42,11 @@ 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) + s = torch.empty(M // block_size, 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, @@ -57,3 +58,218 @@ def triton_block_quant(x, num_stages=6, num_warps=8) return y, s + + +@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..dff30e7 100644 --- a/linghe/quant/channel.py +++ b/linghe/quant/channel.py @@ -4,6 +4,7 @@ """ from typing import Optional + import torch import triton import triton.language as tl @@ -269,4 +270,3 @@ def channel_quant_update(y, x): 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..cc9a4dc 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, flatten=True): 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,9 +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, +def triton_group_quant(x, dtype=torch.float8_e4m3fn, group_size=128, round_scale=False): """ groupwise quantize x, group is in under rowwise format @@ -46,21 +51,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) + 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..ce10a48 100644 --- a/linghe/quant/hadamard.py +++ b/linghe/quant/hadamard.py @@ -136,8 +136,8 @@ 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]( @@ -167,4 +167,4 @@ def triton_hadamard_quant(x, hm): 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/smooth.py b/linghe/quant/smooth.py index 93ac054..ebac037 100644 --- a/linghe/quant/smooth.py +++ b/linghe/quant/smooth.py @@ -13,13 +13,13 @@ @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): + M, T, + N: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, + CALIBRATE: 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, :] @@ -78,14 +78,14 @@ def tokenwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, @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): + M, + N, + H: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, + CALIBRATE: 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, :] @@ -148,8 +148,8 @@ def blockwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, def triton_smooth_quant(x, smooth_scale, x_q=None, x_scale=None, - reverse=False, round_scale=False, - calibrate=False): + reverse=False, round_scale=False, + calibrate=False): """""" M, N = x.shape device = x.device @@ -222,18 +222,18 @@ def triton_smooth_quant(x, smooth_scale, x_q=None, x_scale=None, @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): + 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) @@ -287,8 +287,8 @@ def subrow_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, def triton_subrow_smooth_quant(x, smooth_scale, x_q, x_scale, - subrow_scales, offset, size, - reverse=False, round_scale=False): + subrow_scales, offset, size, + reverse=False, round_scale=False): """""" M, N = x_q.shape W = 128 @@ -335,10 +335,10 @@ def triton_subrow_smooth_quant(x, smooth_scale, x_q, x_scale, @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): + 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)) @@ -363,8 +363,8 @@ def depracated_tokenwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, def triton_depracated_tokenwise_smooth_quant(x, smooth_scale, x_q=None, - x_scale=None, reverse=False, - round_scale=False): + x_scale=None, reverse=False, + round_scale=False): """""" # row-wise read, row-wise write M, N = x.shape @@ -606,10 +606,10 @@ def triton_batch_pad_transpose_smooth_quant(x, @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): + 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, :] @@ -678,10 +678,10 @@ def transpose_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, M, N, P, def triton_transpose_smooth_quant(x, - smooth_scale, - reverse=False, - pad=False, - round_scale=False): + smooth_scale, + reverse=False, + pad=False, + round_scale=False): # col-wise read, row-wise write # M should be padded if M % 32 != 0 """""" @@ -717,14 +717,14 @@ def triton_transpose_smooth_quant(x, @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): + 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, :] @@ -806,11 +806,11 @@ 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): + org_quant_scale, + transpose_smooth_scale, + reverse=True, + pad=False, + round_scale=False): """""" assert reverse M, N = x_q.shape @@ -842,7 +842,6 @@ def triton_transpose_rescale_smooth_quant(x_q, org_smooth_scale, return xt_q, x_scale - """ megatron fp8 training steps: step 0: init w smooth scale w_smooth @@ -870,12 +869,13 @@ def triton_transpose_rescale_smooth_quant(x_q, org_smooth_scale, # 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): +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) + x_scale=x_scale, reverse=False, + round_scale=round_scale) if transpose: xt_q = triton_transpose_and_pad(x_q, out=xt_q, pad=pad) @@ -900,13 +900,13 @@ def triton_smooth_quant_gradient(y, 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) + 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) + transpose_smooth_scale, + reverse=True, + pad=pad, + round_scale=round_scale) else: yt_q, yt_scale = None, None @@ -928,41 +928,40 @@ def triton_smooth_quant_weight(w, 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) + 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) + 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_si = (offset - 1) // N + 1 row_ei = (offset + size) // N col_si = offset % N - col_ei = (offset + size ) % 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) + 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) + 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) - + smooth_scale, + w_q, + quant_scale, + subrow_scales, + offset, + size, + reverse=False, + round_scale=round_scale) 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..a112207 100644 --- a/linghe/tools/benchmark.py +++ b/linghe/tools/benchmark.py @@ -47,9 +47,13 @@ def benchmark_func(fn, *args, n_warmup=10, n_repeat=100, ref_flops=None, 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 = '' if ref_flops is not None: diff --git a/linghe/tools/check.py b/linghe/tools/check.py new file mode 100644 index 0000000..98f3ba2 --- /dev/null +++ b/linghe/tools/check.py @@ -0,0 +1,145 @@ +# -*- 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: + 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() + 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 opt: {opt_str} \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..c2bbf2f 100644 --- a/linghe/tools/util.py +++ b/linghe/tools/util.py @@ -4,6 +4,7 @@ """ import math + import torch @@ -21,6 +22,7 @@ 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, @@ -52,7 +54,7 @@ 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, + 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))) @@ -62,23 +64,66 @@ def torch_group_quant(x, B=128, dtype=torch.float8_e4m3fn, round_scale=False): return xq, scale +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_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): + x = x.float() + m, N = x.shape + 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] + 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] + xt_scale = scale.to(torch.float8_e8m0fnu).view(torch.uint8) + + return x_q, x_scale, xt_q, xt_scale + + def torch_smooth_quant(x, smooth_scale, reverse=False, round_scale=False): x = x.float() x_maxs = x.abs().amax(0) @@ -140,7 +185,7 @@ def torch_make_indices(logits, topk=8, bias=-0.01): 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 @@ -273,7 +318,7 @@ def torch_channel_quant_f_and_b(x, w, y): # 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() @@ -359,49 +404,6 @@ 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' dtype = torch.bfloat16 diff --git a/linghe/utils/add.py b/linghe/utils/add.py index 868e1f0..8e84cb6 100644 --- a/linghe/utils/add.py +++ b/linghe/utils/add.py @@ -4,7 +4,6 @@ """ import torch -from typing import Iterable, Optional, Tuple import triton import triton.language as tl @@ -45,7 +44,7 @@ def inplace_add_kernel(x_ptr, y_ptr, M, N, H: tl.constexpr, W: tl.constexpr, rid * H + tl.arange(0, H)[None, :] < M)) -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,6 +55,7 @@ def triton_inplace_add(x: torch.Tensor, y: torch.Tensor, accum : bool = True): Returns: updated x """ + assert x.is_contiguous() and y.is_contiguous() N = x.shape[-1] M = x.numel() // N # M, N = x.shape 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..c91b4d5 --- /dev/null +++ b/linghe/utils/emb.py @@ -0,0 +1,458 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import triton +import triton.language as tl + + +@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)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + + 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.bfloat16, torch.float32) + shape = x.shape + assert len(shape) == 2 + T = 0 if dtype == torch.float32 else 1 + 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)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + + 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.bfloat16, torch.float32) + T = 0 if dtype == torch.float32 else 1 + 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) + + ids = tl.load(id_ptr + bid * L + sid * B + tl.arange(0, B)) + + 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) + if len(shape) == 2: + B, L = ids.shape + BLOCK = 256 + assert L % BLOCK == 0 + T = 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) + else: + L = shape[0] + B = 1 + BLOCK = 256 + assert L % BLOCK == 0 + T = L // BLOCK + counts = torch.empty((L,), dtype=torch.int32, device=device) + unique_ids = torch.empty((L,), dtype=torch.int32, device=device) + unique_counts = torch.empty((T,), dtype=torch.int32, device=device) + accum_counts = torch.zeros((L + 1,), dtype=torch.int32, device=device) + + num_stages = 3 + num_warps = 1 + 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 = 1 + 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) + + 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, + DIM: tl.constexpr, + T: tl.constexpr, + ): + pid = tl.program_id(axis=0).to(tl.int64) + 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)) + else: + grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) + + 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_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) + T = 0 if dtype == torch.float32 else 1 + 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) + DIM = triton.next_power_of_2(dim) + num_stages = 3 + num_warps = 2 + + grid = (B * L,) + embedding_backward_kernel[grid]( + grad_output, + sorted_ids, + sorted_indices, + accum_counts, + g_ptr, + stride_0, + stride_1, + dim, + B, + L, + DIM, + 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..2829605 --- /dev/null +++ b/linghe/utils/gate.py @@ -0,0 +1,259 @@ +import torch +import triton +import triton.language as tl + + +# 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, + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, + SHARE: tl.constexpr, + TRANSPOSE: 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) + + x_offs = ( + pid * 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) + if TRANSPOSE: + g_offs = ( + sid * bs * DIM + + bid * DIM + + 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) + else: + g = tl.load(gate_ptr + x_offs, mask=x_offs_mask).to(tl.float32) + + rms = tl.rsqrt(tl.sum(x * x, axis=1) / d + eps) + + x = (x * rms[:, None]) * weight * tl.sigmoid(g) + + if TRANSPOSE: + tl.store(out_ptr + g_offs, x, mask=tl.arange(0, D)[None, :] < d) + else: + tl.store(out_ptr + x_offs, x, mask=x_offs_mask) + + +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 tensor has been 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(group_size) == group_size) + assert x.is_contiguous() and gate.is_contiguous() and weight.is_contiguous() + wd = weight.shape[0] + share = wd != dim # all groups share the same weight + d = dim // group_size + device = x.device + + D = triton.next_power_of_2(d) + + 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, + out, + eps, + bs, + length, + dim, + d, + D, + group_size, + share, + transpose, + 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, + DIM: tl.constexpr, + d: tl.constexpr, + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, + T: tl.constexpr, + SHARE: tl.constexpr, + TRANSPOSE: 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, + ) + + x_offs = ( + pid * DIM * T + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, + D)[ + None, :] + ) + x_offs_mask = tl.arange(0, D)[None, :] < d + if TRANSPOSE: + 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 + + 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) + if TRANSPOSE: + g = tl.load(grad_output_ptr + offs, offs_mask).to(tl.float32) + gate = tl.load(gate_ptr + offs, offs_mask).to(tl.float32) + else: + g = tl.load(grad_output_ptr + x_offs, mask=x_offs_mask).to( + tl.float32) + gate = tl.load(gate_ptr + x_offs, mask=x_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) + if TRANSPOSE: + tl.store(dg_ptr + offs, dg, mask=offs_mask) + else: + tl.store(dg_ptr + x_offs, dg, mask=x_offs_mask) + + x_offs += DIM + if TRANSPOSE: + offs += DIM * bs + + 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, transpose=True): + if transpose: + length, bs, dim = gate.shape + else: + bs, length, dim = gate.shape + assert dim <= 8192 and triton.next_power_of_2(group_size) == group_size + assert grad_output.is_contiguous() + d = dim // group_size + wd = weight.shape[0] + share = wd != dim # all groups share the same weight + + 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, + dim, + d, + D, + group_size, + T, + share, + transpose, + num_stages=3, + num_warps=8 + ) + dw = tmp_dw.sum(dim=0) + return dx, dg, dw diff --git a/linghe/utils/gather.py b/linghe/utils/gather.py index 642db26..5d7db2e 100644 --- a/linghe/utils/gather.py +++ b/linghe/utils/gather.py @@ -55,7 +55,6 @@ 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 @@ -69,6 +68,7 @@ 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, @@ -109,10 +109,10 @@ def triton_make_row_id_map( @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, :] @@ -144,7 +144,7 @@ def make_row_id_map_and_indices_kernel(map_ptr, count_ptr, row_map_ptr, offs += b * E -def triton_make_row_id_map_and_indices( +def triton_make_row_id_map_and_index( routing_map: torch.Tensor, num_out_tokens: int, multiple_of: int = 1, @@ -159,13 +159,14 @@ 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, + row_id_indices = torch.zeros((num_out_tokens,), dtype=torch.int32, device=routing_map.device) B = triton.cdiv(n_tokens, T) @@ -183,7 +184,7 @@ def triton_make_row_id_map_and_indices( 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, @@ -204,7 +205,6 @@ def triton_make_row_id_map_and_indices( def index_select_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) @@ -227,15 +227,16 @@ 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,) @@ -245,10 +246,12 @@ def triton_index_select(x, indices, scale=None, out=None, scale_out=None): scale, scale_out, indices, - E, T, N, + E, + T, + N, SCALE, num_stages=3, - num_warps=8 + num_warps=4 ) return out, scale_out @@ -334,8 +337,8 @@ def triton_permute_with_mask_map( 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 +355,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 @@ -450,7 +455,7 @@ def batch_smooth_transpose_smooth_permute_kernel(x_ptr, scale_ptr, oss_ptr, 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) @@ -532,7 +537,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 @@ -553,6 +558,7 @@ def triton_batch_transpose_smooth_permute_with_indices(x, 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_smooth_transpose_smooth_permute_kernel[grid]( x, @@ -578,11 +584,16 @@ def triton_batch_transpose_smooth_permute_with_indices(x, @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, + 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) @@ -592,7 +603,7 @@ def smooth_weighted_permute_with_indices_kernel(grads_ptr, smooth_scale = 1.0 / smooth_scale count = tl.load(count_ptr + pid) ei = tl.load(accum_ptr + pid) - si = ei - count + si = (ei - count).to(tl.int64) 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) @@ -641,9 +652,11 @@ def triton_smooth_weighted_permute_with_indices(grads, x_scale: [bs*topk] x_sum: [bs*topk] """ + assert grads.is_contiguous() M, N = grads.shape n_expert, n = smooth_scales.shape assert N == n, f'{N=} {n=}' + assert triton.next_power_of_2(N) == N E = indices.shape[0] device = grads.device if x_q is None: @@ -664,7 +677,8 @@ def triton_smooth_weighted_permute_with_indices(grads, accum_token_count, indices, x_sum, - M, N, + M, + N, reverse, round_scale, num_stages=3, @@ -675,9 +689,13 @@ def triton_smooth_weighted_permute_with_indices(grads, @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, + grads_scale_ptr, + q_ptr, + ss_ptr, + qs_ptr, + count_ptr, + accum_ptr, + index_ptr, N: tl.constexpr, hs: tl.constexpr, REVERSE: tl.constexpr, @@ -745,11 +763,12 @@ 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 @@ -784,10 +803,14 @@ def triton_smooth_permute_with_indices(grad_data, @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, +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, @@ -860,7 +883,9 @@ def triton_smooth_permute_with_mask_map( - output: output tensor - permuted_scale: permuted scale if scale is not None """ + assert inp.is_contiguous() assert row_id_map.shape[1] == num_experts + assert triton.next_power_of_2(hidden_size) == hidden_size output = torch.empty((num_out_tokens, hidden_size), dtype=torch.float8_e4m3fn, device=row_id_map.device) @@ -872,7 +897,7 @@ def triton_smooth_permute_with_mask_map( (num_out_tokens,), dtype=torch.float32, device=inp.device ) - sm = torch.cuda.get_device_properties(inp.device).multi_processor_count + sm = 128 T = triton.cdiv(num_tokens, sm) grid = (num_experts, sm) smooth_permute_with_mask_map_kernel[grid]( @@ -890,3 +915,149 @@ def triton_smooth_permute_with_mask_map( round_scale ) return output, permuted_scale + + +@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, + 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 flex moe + 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 + diff --git a/linghe/utils/loss.py b/linghe/utils/loss.py index 1eae2fa..64c3ab6 100644 --- a/linghe/utils/loss.py +++ b/linghe/utils/loss.py @@ -7,43 +7,46 @@ 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, +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( + mask=i * B + tl.arange(0, B) < N, other=-1e10).to( tl.float32) - max_logit = tl.maximum(max_logit, tl.max(logit)) - sum_exp += tl.sum(tl.exp(logit)) + 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 - 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) 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 +58,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,9 +71,10 @@ 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 @@ -77,31 +82,56 @@ def triton_softmax_cross_entropy_forward(logits, labels): @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): + 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( + 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): + output_grad, + ignore_index=-100, + inplace=False): """ backward of softmax cross entropy loss Args: @@ -109,27 +139,392 @@ 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 + ) + 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..fc4fb67 --- /dev/null +++ b/linghe/utils/mul.py @@ -0,0 +1,158 @@ +# -*- 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..d2c2689 100644 --- a/linghe/utils/norm.py +++ b/linghe/utils/norm.py @@ -1,64 +1,103 @@ -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)[ + offs = 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)[:, None] < M) & ( + tl.arange(0, N) < n) x = tl.load(x_ptr + offs, - mask=pid * W * T + i * W + tl.arange(0, W)[:, None] < M).to( + mask=mask).to( tl.float32) - rms = tl.sqrt(tl.sum(x * x, axis=1) / N + eps) + 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 + 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, + rms, eps, M, T, + n, N, W, + REUSE, num_stages=3, num_warps=4 ) - return out + return out, rms @triton.jit @@ -66,210 +105,247 @@ def rms_norm_backward_kernel( grad_output_ptr, x_ptr, w_ptr, + rms_ptr, dx_ptr, dw_ptr, eps, M, T, + n, N: tl.constexpr, - W: 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)[ + 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 ) - 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)[ + 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 - - - offs = pid * W * T * N + tl.arange(0, 128)[:, None] * N + tl.arange(0, H)[ - None, :] + 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, :] + 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_output_ptr + toffs, tl.trans(x), mask=indices[None, :] < M) + 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)[ + 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, + scale: Optional[ + torch.Tensor] = None, rms: Optional[torch.Tensor] = None, round_scale: bool = False, output_mode: int = 2): @@ -296,23 +372,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,240 +403,113 @@ 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 ) - 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) - + 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, + 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]( - x, - weight, - out, - scale, - transpose_output, - transpose_scale, - rms, - eps, - M, - T, - N, - N//128, - W, - H, - round_scale, - num_stages=3, - num_warps=16 - ) - scale = scale.t().contiguous() + # 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) return out, scale, rms, transpose_output, transpose_scale -# 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 - - -@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 -): - 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): +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, :] @@ -563,19 +517,19 @@ def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, smooth_scale_ptr 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) + 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 = tl.rsqrt(tl.sum(x * x, axis=1) / N + eps) if OUTPUT: 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) @@ -598,6 +552,7 @@ def triton_rms_norm_and_smooth_quant_forward(x, weight, smooth_scale=None, output_rms=False, 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 @@ -641,3 +596,163 @@ def triton_rms_norm_and_smooth_quant_forward(x, weight, smooth_scale=None, maxs = maxs.amax(0) return out, scale, maxs, rms + +@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..89af89f 100644 --- a/linghe/utils/rearange.py +++ b/linghe/utils/rearange.py @@ -9,9 +9,11 @@ @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) @@ -32,10 +34,10 @@ def split_and_cat_kernel(x_ptr, y_ptr, scale_ptr, scale_output_ptr, count_ptr, 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 +48,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 +64,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, diff --git a/linghe/utils/reduce.py b/linghe/utils/reduce.py index 72b3b3d..24eb767 100644 --- a/linghe/utils/reduce.py +++ b/linghe/utils/reduce.py @@ -59,6 +59,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 @@ -97,7 +98,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, flatten=True): 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 +115,191 @@ 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) + 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=4 + 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): + """ + 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=torch.float32) + 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) + bid = tl.program_id(axis=1).to(tl.int64) sm = tl.num_programs(axis=1) - sums = 0.0 + 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)) + 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 * sm) 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 * sm + 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 + sm = 256 tensor_count = len(xs) - sums = torch.empty((tensor_count, sm), device=device, dtype=torch.float32) - B = 4096 + tmp = torch.empty((tensor_count, sm), device=device, + dtype=torch.float64 if high_precision else torch.float32) + B = 128 grid = (tensor_count, sm) - batch_sum_with_ord_kernel[grid]( + batch_norm_kernel[grid]( ptrs, sizes, - sums, + tmp, + DT, B, ord, + high_precision, num_stages=2, - num_warps=4 + 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..e89a5ae 100644 --- a/linghe/utils/rope.py +++ b/linghe/utils/rope.py @@ -21,24 +21,25 @@ def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, B, 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)[ + 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)[ + q_ptr + pid * q_stride + i * L * q_stride + 2 * D * tl.arange(0, + H)[ :, None] + tl.arange( 0, D)[None, :]) @@ -48,15 +49,17 @@ def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, B, 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( + 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)[ + q_ptr + pid * B * q_stride + i * q_stride + D + 2 * D * tl.arange( + 0, + H)[ :, None] + tl.arange( 0, D)[None, :]) @@ -65,33 +68,36 @@ def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, B, 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( + 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)[ + 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) - + 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)[ + 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)[ + k_ptr + pid * k_stride + i * L * k_stride + 2 * D * tl.arange(0, + h)[ :, None] + tl.arange( 0, D)[None, :]) @@ -101,15 +107,17 @@ def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, B, 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( + 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)[ + k_ptr + pid * B * k_stride + i * k_stride + D + 2 * D * tl.arange( + 0, + h)[ :, None] + tl.arange( 0, D)[None, :]) @@ -118,15 +126,17 @@ def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, B, 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( + 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)[ + k_ptr + pid * k_stride + i * L * k_stride + D + 2 * D * tl.arange( + 0, + h)[ :, None] + tl.arange( 0, D)[None, :]) @@ -137,7 +147,7 @@ def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, B, 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 +157,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 @@ -195,7 +208,7 @@ def half_rope_backward_kernel(q_ptr, k_ptr, freqs_ptr, 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 +217,83 @@ 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, :]) + 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)) - q = q * cos + qr * sin + 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, :]) + 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)) - k = k * cos + kr * sin + 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): +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 @@ -310,18 +333,19 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, D: tl.constexpr, d: tl.constexpr, INTERLEAVED: tl.constexpr, - TRANSPOSED: 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 @@ -333,21 +357,27 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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, :]).to(tl.float32) + q1 = tl.load( + q_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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, :]).to(tl.float32) + q1 = tl.load( + q_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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( @@ -367,8 +397,8 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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)) + 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) k_ptr = qkv_ptr + DD * w @@ -377,22 +407,27 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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, :]).to(tl.float32) k1 = tl.load( k_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + None] + tl.arange( + 0, D)[None, :]).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, :]).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) + None] + tl.arange( + 0, D)[None, :]).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 tl.store( @@ -420,39 +455,240 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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, :]) + v0 = tl.load( + v_ptr + pid * B * stride + i * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + v1 = tl.load( + v_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) else: - v0 = tl.load(v_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + v0 = tl.load( + v_ptr + i * L * stride + pid * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + v1 = tl.load( + v_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + if SILU: + v0 = v0 * tl.sigmoid(v0) + v1 = v1 * tl.sigmoid(v1) + tl.store( vo_ptr + pid * h * DD + i * L * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, + None] + tl.arange(0, + D)[ + None, :], v0) + tl.store( + vo_ptr + pid * h * DD + i * L * h * DD + D + DD * tl.arange(0, h)[:, + None] + tl.arange( + 0, D)[None, :], v1) + + +@triton.jit +def compatible_qk_norm_and_half_rop_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)).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)).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 # 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_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_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, :], 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, :], 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_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_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_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_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_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, :], 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, :], 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, :], 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, :], 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_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_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_p)[:, + None] + tl.arange(0, D)[ - None, :], v0) - + None, :], k0, + mask=k_mask + ) + + if INTERLEAVED: + # 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_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, :], 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, :]) + 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, :]) + 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)[:, + 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) + 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): - + 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 +706,27 @@ 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,33 +736,62 @@ def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, num_stages = 5 num_warps = 2 grid = (L,) - qk_norm_and_half_rope_forward_kernel[grid]( - qkv, - q_norm_weight, k_norm_weight, - freqs, - qo, ko, vo, - B, - stride, - eps, - H, - h, - D // 2, - D // 4, - interleaved, - transposed, - num_stages=num_stages, - num_warps=num_warps - ) + + H_p = triton.next_power_of_2(H) + h_p = triton.next_power_of_2(h) + + if H_p == H and h_p == h: + qk_norm_and_half_rope_forward_kernel[grid]( + qkv, + q_norm_weight, k_norm_weight, + freqs, + qo, ko, vo, + B, + stride, + eps, + H, + h, + D // 2, + D // 4, + interleaved, + transposed, + silu, + num_stages=num_stages, + num_warps=num_warps + ) + else: + compatible_qk_norm_and_half_rop_forward_kernel[grid]( + qkv, + q_norm_weight, k_norm_weight, + freqs, + 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 + ) 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, + q_norm_weight_ptr, + k_norm_weight_ptr, freqs_ptr, dqkv_ptr, - dqw_ptr, dkw_ptr, + dqw_ptr, + dkw_ptr, B, stride, grad_stride, @@ -526,20 +801,21 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, D: tl.constexpr, d: tl.constexpr, INTERLEAVED: tl.constexpr, - TRANSPOSED: 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 - q_w0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)) - q_w1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)) + 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) @@ -556,11 +832,12 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, gq_ptr + i * L * H * DD + pid * H * DD + DD * tl.arange(0, H)[:, None] + tl.arange(0, D)[ - None, :]) + None, :]).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, :]) + 0, D)[None, :]).to(tl.float32) gq_r = tl.reshape(tl.permute( tl.flip(tl.permute(tl.reshape(gq_0, (H, 2, d)), (0, 2, 1)), @@ -568,53 +845,85 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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, :]).to(tl.float32) q1 = tl.load( q_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + None] + tl.arange( + 0, D)[None, :]).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, :]).to(tl.float32) q1 = tl.load( q_ptr + pid * stride + i * L * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + None] + tl.arange( + 0, D)[None, :]).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) 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(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) @@ -632,11 +941,12 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, gk_ptr + i * L * h * DD + pid * h * DD + DD * tl.arange(0, h)[:, None] + tl.arange(0, D)[ - None, :]) + None, :]).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, :]) + 0, D)[None, :]).to(tl.float32) gk_r = tl.reshape(tl.permute( tl.flip(tl.permute(tl.reshape(gk_0, (h, 2, d)), (0, 2, 1)), @@ -644,120 +954,545 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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, :]).to(tl.float32) k1 = tl.load( k_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + None] + tl.arange( + 0, D)[None, :]).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, :]).to(tl.float32) k1 = tl.load( k_ptr + pid * stride + i * L * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]) + None] + tl.arange( + 0, D)[None, :]).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) + tl.store( + dk_ptr + pid * B * grad_stride + i * grad_stride + D + DD * row_offs[ + :, + None] + tl.arange( + 0, D)[None, :], dk_1) 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) + 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(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) + v_ptr = qkv_ptr + DD * w + DD dv_ptr = dqkv_ptr + DD * w + DD else: row_offs = tl.arange(0, 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_0 = tl.load( gv_ptr + i * L * h * DD + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, + 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( + None, :]).to( + tl.float32) + gv_1 = tl.load( gv_ptr + i * L * h * DD + pid * h * DD + D + DD * tl.arange(0, h)[:, None] + tl.arange( - 0, D)[None, :]) + 0, D)[None, :]).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, :]).to(tl.float32) + v1 = tl.load( + v_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + else: + v0 = tl.load( + v_ptr + i * L * stride + pid * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + v1 = tl.load( + v_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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) + tl.store( + dv_ptr + pid * B * grad_stride + i * grad_stride + D + DD * row_offs[ + :, + None] + tl.arange( + 0, D)[None, :], dv_1) else: - tl.store(dv_ptr + pid * grad_stride + i * L * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], v1) + tl.store( + dv_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[ + :, + None] + tl.arange( + 0, D)[None, :], dv_0) + tl.store( + dv_ptr + pid * grad_stride + i * L * grad_stride + D + DD * row_offs[ + :, + None] + tl.arange( + 0, D)[None, :], dv_1) -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): - """ - 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] - transposed: whether qkv is tranposed - transposed: [S, B, dim] - non-transposed: [B, S, dim] +@triton.jit +def compatible_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 - Returns: - - dqkv: gradient of qkv - - dqw: gradient of q_norm_weight - - dkw: gradient of k_norm_weight + 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_w0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)) + q_w1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)) + + 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, 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_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_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_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, :]) + # 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 + pid * stride + i * L * stride + DD * row_offs[:,None] + tl.arange(0, D)[None, :]) + # q1 = tl.load(q_ptr + pid * stride + i * L * stride + D + 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, :], 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] + + 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, 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, 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)).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_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_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_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_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, :], 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, :], 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, :], 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, :], 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] + + if TRANSPOSED: + 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, 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_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_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): + + 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 + 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 + 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: + 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] + 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 +1505,1583 @@ 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, - qkv, - q_norm_weight, k_norm_weight, - freqs, - dqkv, - tmp_dqw, tmp_dkw, - B, - stride, - grad_stride, - eps, - H, - h, - D // 2, - D // 4, - interleaved, - transposed, - num_stages=num_stages, - num_warps=num_warps - ) - dqw = tmp_dqw.sum(0).to(dtype) - dkw = tmp_dkw.sum(0).to(dtype) + if H == H_p and h == h_p: + qk_norm_and_half_rope_backward_kernel[grid]( + gq, + gk, + gv, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + dqkv, + tmp_dqw, + tmp_dkw, + B, + stride, + grad_stride, + eps, + H, + h, + D // 2, + D // 4, + interleaved, + transposed, + silu, + num_stages=num_stages, + num_warps=num_warps + ) + + else: + compatible_qk_norm_and_half_rope_backward_kernel[grid]( + gq, gk, gv, + qkv, + q_norm_weight, k_norm_weight, + freqs, + dqkv, + 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, 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, + 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, H) + tl.arange(0, H) // w * 2 + else: + row_offs = tl.arange(0, H) + + q0 = tl.load(q_ptr + pid * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + q1 = tl.load(q_ptr + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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, H)[:, + None] + tl.arange( + 0, D)[None, :], q1) + + 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)) + q0 = q0 * cos + qr * sin + tl.store( + qo_ptr + pid * 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)).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, h) * (w + 2) + k_ptr = qkv_ptr + DD * w + else: + row_offs = tl.arange(0, h) + k_ptr = qkv_ptr + DD * H + + k0 = tl.load(k_ptr + pid * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + k1 = tl.load( + k_ptr + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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 + tl.store( + ko_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, + None] + tl.arange( + 0, D)[None, :], k1) + + 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)) + k0 = k0 * cos + kr * sin + tl.store( + ko_ptr + pid * h * DD + DD * tl.arange(0, h)[:, + None] + tl.arange(0, + D)[ + None, :], k0) + + if INTERLEAVED: + row_offs = tl.arange(0, h) * (w + 2) + v_ptr = qkv_ptr + DD * w + DD + else: + row_offs = tl.arange(0, 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, :]).to(tl.float32) + v1 = tl.load( + v_ptr + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + + if SILU: + v0 = v0 * tl.sigmoid(v0) + v1 = v1 * tl.sigmoid(v1) + + tl.store( + vo_ptr + pid * h * DD + DD * tl.arange(0, h)[:, + None] + tl.arange(0, + D)[ + None, :], v0) + tl.store( + vo_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, + None] + tl.arange( + 0, D)[None, :], v1) + + +@triton.jit +def compatible_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 + stride = qkv.stride(0) # qkv may be a slice of a tensor + D = Dim // (H + 2 * h) + B = cu_seqlens_q.size(0) - 1 + PB = max(triton.next_power_of_2(B), 128) # reduce jit + 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) + + if PH == H and ph == 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, + D // 2, + D // 4, + interleaved, + silu, + cp_size, + reuse, + num_stages=num_stages, + num_warps=num_warps + ) + else: + compatible_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, + 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, H) + tl.arange(0, H) // w * 2 + else: + row_offs = tl.arange(0, H) + + gq_0 = tl.load( + gq_ptr + pid * H * DD + DD * tl.arange(0, H)[:, + None] + tl.arange(0, + D)[ + None, :]).to(tl.float32) + gq_1 = tl.load( + gq_ptr + pid * H * DD + D + DD * tl.arange(0, H)[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + + 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_0 = gq_0 * cos + gq_r * sin + + q0 = tl.load(q_ptr + pid * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + q1 = tl.load( + q_ptr + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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) + tl.store(dq_ptr + pid * grad_stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :], dq_1) + + 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, h) * (w + 2) + k_ptr = qkv_ptr + DD * w + dk_ptr = dqkv_ptr + DD * w + else: + row_offs = tl.arange(0, 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, h)[:, + None] + tl.arange(0, + D)[ + None, :]).to(tl.float32) + gk_1 = tl.load( + gk_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + + 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_0 = gk_0 * cos + gk_r * sin + + k0 = tl.load(k_ptr + pid * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + k1 = tl.load( + k_ptr + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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) + tl.store(dk_ptr + pid * grad_stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :], dk_1) + + 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, h) * (w + 2) + v_ptr = qkv_ptr + DD * w + DD + dv_ptr = dqkv_ptr + DD * w + DD + else: + row_offs = tl.arange(0, 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, h)[:, + None] + tl.arange(0, + D)[ + None, :]).to(tl.float32) + gv_1 = tl.load( + gv_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + + if SILU: + v0 = tl.load(v_ptr + pid * stride + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).to(tl.float32) + v1 = tl.load( + v_ptr + pid * stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :]).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) + tl.store(dv_ptr + pid * grad_stride + D + DD * row_offs[:, + None] + tl.arange( + 0, D)[None, :], dv_1) + + +@triton.jit +def compatible_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 = max(triton.next_power_of_2(B), 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) + + if PH == H and ph == 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, + D // 2, + D // 4, + interleaved, + silu, + cp_size, + reuse, + num_stages=num_stages, + num_warps=num_warps + ) + else: + compatible_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 + ) + 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 = max(triton.next_power_of_2(B), 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 = max(triton.next_power_of_2(B), 128) + assert B <= 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..8f8ada6 100644 --- a/linghe/utils/scatter.py +++ b/linghe/utils/scatter.py @@ -4,12 +4,12 @@ """ 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, @@ -45,12 +45,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 @@ -70,7 +71,7 @@ def triton_aligned_scatter_add(x: torch.Tensor, # 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 +82,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 +107,14 @@ 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) - sm = torch.cuda.get_device_properties(x.device).multi_processor_count + sm = 512 T = triton.cdiv(M, sm) num_stages = 5 @@ -140,37 +143,49 @@ def triton_scatter_add(x, outputs, indices): @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)) 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) - 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( @@ -189,16 +204,20 @@ 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, + output = torch.empty((num_tokens, n), dtype=grad.dtype, device="cuda") PROB = probs is not None if PROB: + assert probs.is_contiguous() restore_probs = torch.zeros((num_tokens, num_experts), - dtype=probs.dtype, device="cuda") + dtype=probs.dtype, + device="cuda") else: restore_probs = None @@ -212,8 +231,9 @@ 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 diff --git a/linghe/utils/silu.py b/linghe/utils/silu.py index 90545b3..7b67dd1 100644 --- a/linghe/utils/silu.py +++ b/linghe/utils/silu.py @@ -4,41 +4,100 @@ """ 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, +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 + + 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) + + +@triton.jit +def weighted_silu_forward_kernel(x_ptr, weight_ptr, out_ptr, M, + N, + H: tl.constexpr, W: tl.constexpr, WEIGHT: tl.constexpr): - pid = tl.program_id(axis=0) - - row_offs = pid * W * T * n + tl.arange(0, W)[:, None] * n - col_offs = tl.arange(0, n)[None, :] + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + n = (N // 2).to(tl.int64) - 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 + 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,68 +106,88 @@ 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, +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) - - 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) + n = (N // 2).to(tl.int64) + + 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, @@ -125,29 +204,31 @@ 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, T, + M, N, - N // 2, + H, W, WEIGHT, num_stages=3, @@ -163,52 +244,51 @@ def silu_and_block_quant_forward_kernel(x_ptr, transpose_scale_ptr, M, n: tl.constexpr, + H: tl.constexpr, + W: 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 * 2 + cid * 128 + tl.arange(0, 128)[:, - None] * n * 2 + tl.arange(0, 128)[ - None, :] - indices = rid * 128 + tl.arange(0, 128) + 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 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 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(scale_ptr + rid * 128 + cid * M + tl.arange(0, 128), 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 * 128 * n + cid * 128 + tl.arange(0, 128)[:, - None] * n + tl.arange(0, - 128)[ - None, :], xq, + 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: 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), + 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 * 128 + cid * 128 * M + tl.arange(0, - 128)[ - :, - None] * M + tl.arange( - 0, 128)[ - None, - :], + xq = (x / scale).to(transpose_output_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) @@ -233,21 +313,28 @@ def triton_silu_and_block_quant_forward(x, - 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 // 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, + scale = torch.empty((n // 128, M), device=device, dtype=torch.float32) - transpose_output = torch.empty((N // 2, M), device=device, + transpose_output = torch.empty((n, M), device=device, dtype=torch.float8_e4m3fn) - transpose_scale = torch.empty((triton.cdiv(M, 128), N // 2), device=device, + transpose_scale = torch.empty((triton.cdiv(M, 128), n), device=device, dtype=torch.float32) - - grid = (triton.cdiv(M, 128), n // 128) + if output_mode == 0: + H, W, num_warps = 64, 128, 4 + elif output_mode == 1: + H, W, num_warps = 128, 64, 4 + else: + H, W, num_warps = 128, 128, 8 + assert n % W == 0 + grid = (triton.cdiv(M, H), n // W) silu_and_block_quant_forward_kernel[grid]( x, out, @@ -256,10 +343,12 @@ def triton_silu_and_block_quant_forward(x, transpose_scale, M, n, + H, + W, round_scale, output_mode, num_stages=2, - num_warps=16 + num_warps=num_warps ) return out, scale, transpose_output, transpose_scale @@ -285,14 +374,14 @@ def silu_and_block_quant_backward_kernel(g_ptr, x_ptr, 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) + 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) - sigmoid = tl.sigmoid(x1.to(tl.float32)) + tl.arange(0, 128)[None, :], mask=mask).to(tl.float32) + sigmoid = tl.sigmoid(x1) dx1 = sigmoid * g * x2 * ( - 1 + x1 * (1 - sigmoid)) # change order to trigger autocast + 1 + x1 * (1 - sigmoid)) scale1 = tl.maximum( tl.max(dx1.abs(), 1) / 448, 1e-30) if ROUND: @@ -311,7 +400,7 @@ def silu_and_block_quant_backward_kernel(g_ptr, x_ptr, transpose_dx_scale_ptr + rid * n * 2 + cid * 128 + tl.arange(0, 128), scale1) - qdx1 = (dx1 / scale1[None, :]).to(dx_ptr.dtype.element_ty) + qdx1 = (dx1 / scale1[None, :]).to(transpose_dx_ptr.dtype.element_ty) tl.store(transpose_dx_ptr + toffs, tl.trans(qdx1), mask=idx[None, :] < M) dx2 = sigmoid * g * x1 @@ -333,7 +422,7 @@ def silu_and_block_quant_backward_kernel(g_ptr, x_ptr, 128), scale2) - qdx2 = (dx2 / scale2[None, :]).to(dx_ptr.dtype.element_ty) + 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) @@ -354,6 +443,7 @@ def triton_silu_and_block_quant_backward(g, x, - 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 @@ -365,7 +455,7 @@ def triton_silu_and_block_quant_backward(g, x, transpose_dx_scale = torch.empty(scale_shape, device=device, dtype=torch.float32) - assert M % 128 == 0 + assert M % 128 == 0 and N % 256 == 0 grid = (M // 128, N // 256) silu_and_block_quant_backward_kernel[grid]( g, @@ -391,10 +481,9 @@ def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, transpose_scale_ptr, count_ptr, accum_ptr, - n: tl.constexpr, + n, E: tl.constexpr, - ROUND: tl.constexpr, - OUTPUT_MODE: tl.constexpr): + ROUND: tl.constexpr): eid = tl.program_id(axis=0) rid = tl.program_id(axis=1) cid = tl.program_id(axis=2) @@ -407,6 +496,7 @@ def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, if rid >= c: return + n = n.to(tl.int64) nb = n // 128 counts = tl.load(count_ptr + tl.arange(0, E)) @@ -433,30 +523,235 @@ def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, 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))) + tl.store( + scale_ptr + si * nb + cid * count + rid * 128 + tl.arange(0, 128), + scale, mask=indices < count) + + xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) + tl.store(out_ptr + hoffs, xq, mask=mask) + + 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 * 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) + + +@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, + n, + B: 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) + 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) + 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) + + 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 * 128 + tl.arange(0, 128), + 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 + + # transpose + counts = tl.load(count_ptr + tl.arange(0, E)) + n_blocks = tl.cdiv(counts, 128) + transpose_soff = 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, B)[None, :] + toffs = si * n + rid * 128 + cid * count * 128 + tl.arange(0, B)[:, + None] * count + tl.arange( + 0, 128)[ + None, :] + tsoffs = transpose_soff * n + rid * n + cid * 128 + tl.arange( + 0, B) + indices = rid * 128 + tl.arange(0, 128) + 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) + + x = x1 * tl.sigmoid(x1) * x2 * w[:, None] - 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))) tl.store( - transpose_scale_ptr + transpose_scale_off * n + rid * n + cid * 128 + tl.arange( - 0, 128), scale) + transpose_scale_ptr + tsoffs, scale) + + xq = tl.trans((x / scale).to(transpose_output_ptr.dtype.element_ty)) - xq = tl.trans((x / scale).to(out_ptr.dtype.element_ty)) tl.store(transpose_output_ptr + toffs, xq, mask=indices[None, :] < count) + offs += B + toffs += count * B + tsoffs += 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, + n, + B: 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) + 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) + + 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, + n, + B: 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) + 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) + + 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, @@ -485,6 +780,8 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, - 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] @@ -492,41 +789,123 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, 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) - 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, + transpose_output = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) - transpose_scale = torch.empty((blocks * n), device=device, + 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) if M == 0: return out, scale, transpose_output, transpose_scale accums = torch.cumsum(counts, 0) - grid = (n_experts, triton.cdiv(max(splits), 128), n // 128) - batch_weighted_silu_and_block_quant_forward_kernel[grid]( - x, - weight, - out, - scale, - transpose_output, - transpose_scale, - counts, - accums, - n, - len(splits), - round_scale, - output_mode, - num_stages=2, - num_warps=8 - ) + 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, + n, + B, + len(splits), + 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, + n, + B, + len(splits), + round_scale, + num_stages=2, + num_warps=2 + ) + else: + grid = (n_experts, triton.cdiv(max(splits), 128), n // 128) + batch_weighted_silu_and_block_quant_forward_kernel[grid]( + x, + weight, + out, + scale, + transpose_output, + transpose_scale, + counts, + accums, + n, + len(splits), + round_scale, + num_stages=2, + num_warps=8 + ) + # 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, + # n, + # B, + # len(splits), + # round_scale, + # num_stages=2, + # num_warps=8 + # ) + + # 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, + # n, + # B, + # len(splits), + # round_scale, + # num_stages=2, + # num_warps=2 + # ) + # 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, + # n, + # B, + # len(splits), + # round_scale, + # num_stages=2, + # num_warps=2 + # ) return out, scale, transpose_output, transpose_scale @@ -540,7 +919,7 @@ def batch_weighted_silu_and_block_quant_backward_kernel(g_ptr, x_ptr, transpose_dx_ptr, transpose_dx_scale_ptr, dw_ptr, - n: tl.constexpr, + n, E: tl.constexpr, ROUND: tl.constexpr): eid = tl.program_id(axis=0) @@ -549,10 +928,14 @@ 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 + # 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)) @@ -565,13 +948,13 @@ def batch_weighted_silu_and_block_quant_backward_kernel(g_ptr, x_ptr, 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) + 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)) + mask=idx[:, None] < count).to(tl.float32) + sigmoid = tl.sigmoid(x1) dw = tl.sum(sigmoid * x1 * x2 * g, 1) tl.store(dw_ptr + si * nb + cid + idx * nb, dw, mask=idx < count) @@ -627,6 +1010,146 @@ def batch_weighted_silu_and_block_quant_backward_kernel(g_ptr, x_ptr, 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, + n, + B: 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) + 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) + 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) + sigmoid = tl.sigmoid(x1) + + dw = tl.sum(sigmoid * x1 * 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 * B + tl.arange(0, + B), + scale, mask=idx < count) + + tl.store(dx_ptr + offs, dx / scale[:, None], 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))) + 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, + n, + B: 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) + 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) + 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) + sigmoid = tl.sigmoid(x1) + + dx = sigmoid * g * x2 * w * (1 + x1 * (1 - sigmoid)) + + 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 = sigmoid * g * x1 * w + + 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, @@ -648,11 +1171,11 @@ 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 n = N // 2 n_expert = counts.shape[0] - assert N <= 8192 and 8192 % N == 0 + assert n % 128 == 0 + assert g.is_contiguous() assert splits is not None, 'batch mode need splits to launch kernels' device = x.device @@ -660,24 +1183,42 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, 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_scale = torch.empty((M, N // 128), device=device, dtype=torch.float32) s = sum([(x + 127) // 128 for x in splits]) - transpose_dx = torch.empty((N * M), device=device, + transpose_dx = torch.empty((M, N), device=device, dtype=torch.float8_e4m3fn) - transpose_dx_scale = torch.empty((s * N), device=device, + 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 - # grid = (n_expert, triton.cdiv(max(splits), 128)) - grid = (n_expert, triton.cdiv(max(splits), 128), N // 256) + # 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, + # n, + # n_expert, + # 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_kernel[grid]( + batch_weighted_silu_and_block_quant_backward_n_kernel[grid]( g, x, weight, @@ -685,28 +1226,45 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, accums, dx, dx_scale, - transpose_dx, - transpose_dx_scale, dws, n, + B, n_expert, round_scale, - num_stages=3, - num_warps=16 + num_stages=2, + num_warps=4 ) dw = dws.sum(1, keepdim=True).to(weight.dtype) - return dx, dx_scale, dw, transpose_dx, transpose_dx_scale + 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, + n, + B, + n_expert, + round_scale, + num_stages=2, + num_warps=4 + ) + 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, + max_ptr, M, T, n: tl.constexpr, + W: tl.constexpr, ROUND: tl.constexpr, + CALIBRATE: tl.constexpr): pid = tl.program_id(axis=0) row_offs = pid * T * W * n + tl.arange(0, W)[:, None] * n @@ -722,7 +1280,7 @@ def silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out_ptr, scale 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 + x = x1 * tl.sigmoid(x1) * x2 if CALIBRATE: maxs = tl.maximum(x.abs(), maxs) x = x * smooth_scale @@ -741,12 +1299,14 @@ def silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out_ptr, scale # 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, max_ptr, M, + T: tl.constexpr, + n: tl.constexpr, + B: tl.constexpr, + ROUND: tl.constexpr, + CALIBRATE: tl.constexpr): pid = tl.program_id(axis=0) # rowwise read with block size [T, B] @@ -760,7 +1320,7 @@ def compatible_silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out 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 + x = x1 * tl.sigmoid(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) @@ -779,7 +1339,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 +1347,13 @@ 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=None, out=None, + scale=None, + maxs=None, round_scale=False, + calibrate=False): """""" + assert x.is_contiguous() M, N = x.shape n = N // 2 device = x.device @@ -803,11 +1363,10 @@ 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) @@ -853,26 +1412,22 @@ def triton_silu_and_smooth_quant_forward(x, smooth_scale=None, out=None, scale=N if calibrate: maxs = maxs.amax(0) - return out, scale, maxs - - - @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): + 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)[ @@ -881,8 +1436,9 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, :] 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,12 +1449,12 @@ 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)) @@ -908,17 +1464,22 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, 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))) - tl.store(transpose_dx_scale_ptr + pid * n * 2 + i * B + tl.arange(0, B), 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))) + 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 @@ -954,7 +1515,7 @@ 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) dx1 = g * x2 * sigmoid * ( 1 + x1 * (1 - sigmoid)) * smooth_scale_1 dx2 = g * x1 * sigmoid * smooth_scale_2 @@ -967,6 +1528,7 @@ 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, @@ -977,37 +1539,42 @@ def _requant_kernel(x_ptr, scale_ptr, scales_ptr, ): 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): + 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) + transpose_dx_scale = torch.empty(scale_shape, device=device, + dtype=torch.float32) T = 32 B = 32 assert M % T == 0 and n % B == 0 - transpose_dx_scales = torch.empty((M // T, N), device=device, dtype=torch.float32) + transpose_dx_scales = torch.empty((M // T, N), device=device, + dtype=torch.float32) grid = (M // T,) silu_and_smooth_quant_backward_kernel[grid]( g, @@ -1030,10 +1597,10 @@ def triton_silu_and_smooth_quant_backward(g, x, 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) + N, + M, + B, + T) return dx, dx_scale, transpose_dx, transpose_dx_scale @@ -1044,7 +1611,8 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, out_ptr, scale_ptr, max_ptr, count_ptr, - accum_ptr, M, + accum_ptr, + M, n: tl.constexpr, W: tl.constexpr, ROUND: tl.constexpr, @@ -1056,7 +1624,7 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, 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 @@ -1078,7 +1646,7 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, w = tl.load(weight_ptr + si + indices, mask=indices < count).to( tl.float32)[:, None] - x = x1 / (1 + tl.exp(-x1)) * x2 + x = x1 * tl.sigmoid(x1) * x2 if CALIBRATE: maxs = tl.maximum(x.abs(), maxs) @@ -1109,15 +1677,16 @@ def triton_batch_weighted_silu_and_smooth_quant_forward(x, reverse=False, calibrate=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 + sm = 128 tmp_maxs = None if scale is None: scale = torch.empty((M,), device=device, dtype=torch.float32) @@ -1187,7 +1756,7 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, 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 @@ -1235,7 +1804,7 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, 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)) + sigmoid = tl.sigmoid(x1) dx1 = g * x2 * sigmoid * ( 1 + x1 * (1 - sigmoid)) * w dx2 = g * x1 * sigmoid * w @@ -1301,7 +1870,7 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, 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)) + sigmoid = tl.sigmoid(x1) dx1 = g * x2 * sigmoid * ( 1 + x1 * (1 - sigmoid)) * smooth_scale_1 * w dx2 = g * x1 * sigmoid * smooth_scale_2 * w @@ -1360,6 +1929,7 @@ def triton_batch_weighted_silu_and_smooth_quant_backward(g, x, weight, reverse=True, round_scale=False): """""" + assert g.is_contiguous() assert round_scale M, N = x.shape n = N // 2 @@ -1428,4 +1998,3 @@ def triton_batch_weighted_silu_and_smooth_quant_backward(g, x, weight, 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..8848b71 --- /dev/null +++ b/linghe/utils/topk.py @@ -0,0 +1,289 @@ +# -*- 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)) + + 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) + + +def triton_topk_forward(x, k, dim=-1): + """ + calculate topk. + Args: + x: input tensor. + k: topk + Returns: + values: topk values + indices: topk indices + """ + device = x.device + shape = x.shape + assert dim == -1 and len(shape) <= 3 + assert x.is_contiguous() + if len(shape) == 3: + M, B, N = shape + g = M * B + values = torch.empty((M, B, k), device=device, dtype=x.dtype) + indices = torch.empty((M, B, k), device=device, dtype=torch.int64) + else: + M, N = shape + g = M + values = torch.empty((M, k), device=device, dtype=x.dtype) + indices = torch.empty((M, k), device=device, dtype=torch.int64) + grid = (g,) + topk_forward_kernel[grid]( + x, + values, + indices, + N, + k, + num_stages=2, + num_warps=2 + ) + 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 and len(shape) <= 3 + assert grad_output.is_contiguous() + if len(shape) == 3: + M, B, k = shape + g = M * B + dx = torch.zeros((M, B, N), device=device, dtype=grad_output.dtype) + else: + M, k = shape + g = M + dx = torch.zeros((M, 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 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) + + +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,) + group_topk_score_forward_kernel[grid]( + x, + expert_bias, + probs, + routing_map, + scaling_factor, + eps, + N, + k, + num_groups, + group_topk, + BIAS, + num_stages=1, + 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..b99ad35 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 @@ -41,8 +41,8 @@ def transpose_kernel(x_ptr, t_ptr, M, N, H: tl.constexpr, W: tl.constexpr, @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 +51,40 @@ 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: + 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)) + + +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 @@ -86,7 +111,7 @@ def triton_transpose(x: torch.Tensor, 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] @@ -101,18 +126,41 @@ 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 @@ -144,7 +192,6 @@ def transpose_and_pad_kernel(x_ptr, t_ptr, mask=(rid * H + tl.arange(0, H)[None, :] < P)) - def triton_transpose_and_pad(x, out=None, pad=True): """ transpose x and padding the column size to be mutiplier of 32, @@ -158,6 +205,7 @@ 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 device = x.device @@ -206,15 +254,20 @@ 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, + xts = torch.empty((M * n_experts, N), + device=device, dtype=xs[0].dtype) - pointers = torch.tensor([x.data_ptr() for x in xs], device=xs[0].device) - H = 64 + 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]( @@ -224,11 +277,7 @@ def triton_batch_transpose(xs, xts=None): 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 @@ -268,6 +317,7 @@ def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): Returns: x_t: output tensor """ + assert x.is_contiguous() assert pad # block shape:[H,W] M, N = x.shape @@ -331,6 +381,7 @@ 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 diff --git a/linghe/utils/unary.py b/linghe/utils/unary.py index d3dd682..14343d7 100644 --- a/linghe/utils/unary.py +++ b/linghe/utils/unary.py @@ -9,7 +9,7 @@ @triton.jit -def calculate_smooth_scale_kernel(x_ptr, y_ptr, min_value, smooth_coef, +def calculate_smooth_scale_kernel(x_ptr, y_ptr, min_value, smooth_coef, N, B: tl.constexpr, EVEN: tl.constexpr, @@ -17,23 +17,25 @@ def calculate_smooth_scale_kernel(x_ptr, y_ptr, min_value, smooth_coef, 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)) + 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) + 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 + 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/dev.py b/scripts/dev.py new file mode 100644 index 0000000..b46d113 --- /dev/null +++ b/scripts/dev.py @@ -0,0 +1,30 @@ +import torch +import triton +import triton.language as tl + + +def test_cpu_gpu_diff(): + x = torch.randn((128, 128), dtype=torch.float32) * 100 + cos = torch.cos(x) + c = torch.cos(x.cuda()) + torch.testing.assert_close(cos, c.cpu()) + + +@triton.jit +def index_overflow(x): + i = tl.program_id(0) + ptr = x + i * 2 ** 30 + offs = 2 ** 30 + 2 ** 30 + 2 ** 30 + 2 ** 30 + # offs = (i).to(tl.int64) * 2 ** 30 + 5*2**32 + tl.store(x + i, offs) + + +def test_index_overflow(): + x = torch.zeros((128,), dtype=torch.int64, device='cuda:0') + index_overflow[(128,)](x) + print(x) + + +if __name__ == '__main__': + # test_cpu_gpu_diff() + test_index_overflow() diff --git a/scripts/plot_input_output.py b/scripts/plot_input_output.py index 4d0395c..1cc5045 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] + 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 + 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) + 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) + 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 + 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] + 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 + N, K = d['w'].shape + M = d['x'].numel() // K - xq = d['x'].float().view(M,K) + xq = d['x'].float().view(M, K) xs = d['xs'].float() xm = d['x_smooth_scale'] - x = xq*xm*xs[:,None] + x = xq * xm * xs[:, None] wq = d['w'].float() ws = d['ws'].float() wm = d['w_smooth_scale'] - w = wq*wm*ws[:,None] + 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] + dy = dyq / dym * dys[:, None] dytq = d['dyt'].float().t() dyts = d['dyts'] - dytm=d['dyt_smooth_scale'] - dyt = dytq/dytm*dyts[:,None] + dytm = d['dyt_smooth_scale'] + dyt = dytq / dytm * dyts[:, None] x = x.cuda() w = w.cuda().transpose().contiguous() @@ -61,34 +60,34 @@ 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) + 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' fig, ax = plt.subplots(figsize=(8, 12)) ax.imshow(xb, cmap='gray') # plt.show() plt.axis('off') - plt.savefig(f"figures/{prefix}_x.{fmt}", bbox_inches='tight',dpi=600) + 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') # plt.show() plt.axis('off') - plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches='tight',dpi=600) + plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches='tight', dpi=600) plt.close('all') fig, ax = plt.subplots(figsize=(8, 12)) @@ -115,26 +114,26 @@ def read_fp8_inputs(prefix='fc2'): # fp8 if False: prefix = 'fc2' - x,w,dy,dyt,xm,wm = read_fp8_inputs(prefix=prefix) + 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' fig, ax = plt.subplots(figsize=(8, 12)) ax.imshow(xb, cmap='gray') # plt.show() plt.axis('off') - plt.savefig(f"figures/{prefix}_x.{fmt}", bbox_inches='tight',dpi=600) + 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') # plt.show() plt.axis('off') - plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches='tight',dpi=600) + plt.savefig(f"figures/{prefix}_w.{fmt}", bbox_inches='tight', dpi=600) plt.close('all') fig, ax = plt.subplots(figsize=(8, 12)) @@ -143,6 +142,3 @@ def read_fp8_inputs(prefix='fc2'): 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..d8c192a 100644 --- a/scripts/reproduce_triton_bug.py +++ b/scripts/reproduce_triton_bug.py @@ -3,24 +3,23 @@ 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, :] @@ -28,18 +27,18 @@ def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, smooth_scale_ptr smooth_scale = 1.0 / tl.maximum(smooth_scale, 1e-30) # triton 3.3.1 has bug with N = 2048 and calibrate=True if CALIBRATE: - maxs = tl.zeros((N, ), dtype=tl.float32) + 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 +53,11 @@ 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): + 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 +69,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,12 +92,11 @@ 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__': @@ -106,16 +104,25 @@ def triton_rms_norm_and_quant_forward(x, weight, smooth_scale, eps=1e-6, N = 2048 dtype = torch.bfloat16 device = 'cuda:0' - calibrate = True - # bug condition: triton=3.3.1 N=2048 calibrate=True num_warps=4 + 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 + # 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) - print(f'{q=}\n{scale=}') \ No newline at end of file + 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..88316f7 100644 --- a/scripts/test.sh +++ b/scripts/test.sh @@ -1,18 +1,30 @@ -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_mxfp8_quant.py" && python test_mxfp8_quant.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..ae32ca6 100644 --- a/setup.py +++ b/setup.py @@ -14,7 +14,7 @@ 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..b56fb8e 100644 --- a/tests/test_add.py +++ b/tests/test_add.py @@ -6,7 +6,7 @@ import torch from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check +from linghe.tools.check import output_check from linghe.utils.add import triton_inplace_add @@ -30,9 +30,8 @@ def test_triton_inplace_add(M=4096, N=4096, bench=False): out_ref = outputs + x output_check(out_ref, out, 'sum') - n_repeat = 100 - if bench: + n_repeat = 100 ref_time = benchmark_func(torch_add, x, out, accum=False, n_repeat=n_repeat) benchmark_func(triton_inplace_add, out, x, accum=False, diff --git a/tests/test_blockwise_fp8_gemm.py b/tests/test_blockwise_fp8_gemm.py index 497b791..9e5d3e6 100644 --- a/tests/test_blockwise_fp8_gemm.py +++ b/tests/test_blockwise_fp8_gemm.py @@ -5,11 +5,10 @@ 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_bb_fp8_gemm, \ + triton_tt_fp8_gemm from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check - +from linghe.tools.check import output_check def test_triton_bb_gemm(M=4096, N=4096, K=4096, bench=False): @@ -26,23 +25,23 @@ 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') + 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.to(dtype), y, name='y', rtol=0.05, atol=1.0) 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) - + 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): @@ -59,27 +58,23 @@ def test_triton_tt_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, 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') + 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.to(dtype), y, 'y', atol=1.0, rtol=0.05) 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) - + 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) - - + test_triton_bb_gemm(M=4096, N=8192, K=2048, bench=False) + test_triton_tt_gemm(M=4096, N=8192, K=2048, bench=False) diff --git a/tests/test_blockwise_quant.py b/tests/test_blockwise_quant.py new file mode 100644 index 0000000..e921e34 --- /dev/null +++ b/tests/test_blockwise_quant.py @@ -0,0 +1,115 @@ +import torch + +from linghe.quant.block import triton_block_quant, triton_blockwise_quant, \ + triton_batch_blockwise_quant +from linghe.tools.benchmark import benchmark_func +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 + + +def test_block_quant(M=8192, N=4096, bench=False): + 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') + + if bench: + benchmark_func(triton_block_quant, x, + round_scale=True, + ref_bytes=M * N * 4) + + +def test_blockwise_quant(M=8192, N=4096, bench=False): + 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') + + if bench: + benchmark_func(triton_blockwise_quant, x, + round_scale=True, + ref_bytes=M * N * 4) + + +def test_batch_block_quant(M=16384, N=2048, n_experts=32, topk=2, bench=False): + 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') + + if bench: + benchmark_func(triton_batch_blockwise_quant, x, token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=M * N * 4) + + +if __name__ == '__main__': + test_block_quant(M=8192, N=4096, bench=False) + test_blockwise_quant(M=8192, N=4096, bench=False) + test_batch_block_quant(M=16384, N=2048, n_experts=32, topk=2, bench=False) diff --git a/tests/test_channel_quant.py b/tests/test_channel_quant.py index 5b4fd49..f08a64b 100644 --- a/tests/test_channel_quant.py +++ b/tests/test_channel_quant.py @@ -6,11 +6,11 @@ import torch from linghe.quant.channel import (triton_deprecated_tokenwise_row_quant, - triton_row_quant, - triton_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) +from linghe.tools.check import output_check +from linghe.tools.util import torch_row_quant def test_row_quant(M=4096, N=4096, round_scale=True, bench=False): @@ -21,12 +21,12 @@ def test_row_quant(M=4096, N=4096, round_scale=True, bench=False): 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') + output_check(x_q_ref, x_q, name='data') + output_check(x_scale_ref, x_scale, name='scale') if bench: ref_time = benchmark_func(torch_row_quant, x, n_repeat=100, diff --git a/tests/test_channelwise_fp8_gemm.py b/tests/test_channelwise_fp8_gemm.py index ce5a195..b5d5490 100644 --- a/tests/test_channelwise_fp8_gemm.py +++ b/tests/test_channelwise_fp8_gemm.py @@ -6,19 +6,18 @@ import torch from linghe.gemm.channelwise_fp8_gemm import triton_scaled_mm - -from linghe.utils.add import triton_inplace_add from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check +from linghe.tools.check import output_check +from linghe.utils.add import triton_inplace_add 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) + 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,10 +26,7 @@ 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): - dtype = torch.bfloat16 device = 'cuda:0' @@ -42,12 +38,12 @@ 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_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') + accum=False) + output_check(y_ref, y, name='y', atol=-1) if bench: y_bf16 = torch.randn(M, N, dtype=dtype, device=device) @@ -57,24 +53,27 @@ def test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False): 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=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(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) + 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) + 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) + 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) - + test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False) diff --git a/tests/test_dist_loss.py b/tests/test_dist_loss.py new file mode 100644 index 0000000..76d519b --- /dev/null +++ b/tests/test_dist_loss.py @@ -0,0 +1,159 @@ +# -*- 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..64ab114 --- /dev/null +++ b/tests/test_embedding.py @@ -0,0 +1,155 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch + +from linghe.facade.emb import (embedding_lookup, + fused_accumulation_embedding_lookup) +from linghe.tools.benchmark import benchmark_func +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 + ) + + +def test_scan(M=4096, bench=False): + device = 'cuda: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') + + if bench: + ref_time = benchmark_func(triton_scan_and_count, sorted_ids) + + +def test_embedding(B=2, M=4096, V=150000, D=4096, transpose=False, bench=False): + 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 + weights.grad = torch.zeros((V, D), dtype=dtype, device=device) + + 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) + grad_ref = weights.grad.clone().detach() + + grad = weights.grad + grad.zero_() + y = triton_embedding_forward(input_ids, weights.data_ptr(), D, dtype) + output_check(y_ref, y, name='y') + + triton_embedding_backward(dy, input_ids, grad.data_ptr(), grad.dtype) + output_check(grad_ref, grad.to(dtype), name='grad') + + grad.zero_() + 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') + + if bench: + ref_bytes = B * M * D * 4 + ref_time = benchmark_func(embedding.forward, input_ids, + ref_bytes=ref_bytes) + benchmark_func(embedding_lookup, input_ids, weights, + ref_time=ref_time, ref_bytes=ref_bytes) + + ref_time = benchmark_func(y_ref.backward, dy, retain_graph=True, + ref_bytes=ref_bytes) + benchmark_func(triton_atomic_embedding_backward, dy, input_ids, + grad.data_ptr(), grad.dtype, + ref_time=ref_time, ref_bytes=ref_bytes) + benchmark_func(triton_sync_embedding_backward, dy, input_ids, + grad.data_ptr(), grad.dtype, + ref_time=ref_time, ref_bytes=ref_bytes) + benchmark_func(triton_embedding_backward, dy, input_ids, + grad.data_ptr(), grad.dtype, + ref_time=ref_time, ref_bytes=ref_bytes) + benchmark_func(y.backward, dy, retain_graph=True, + ref_time=ref_time, ref_bytes=ref_bytes) + + +def test_fused_embedding(B=2, M=4096, V=150000, D=4096, use_main_grad=True, + transpose=False, bench=False): + dtype = torch.bfloat16 + device = 'cuda:0' + grad_name = 'main_grad' if use_main_grad else 'grad' + + 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 + weights.grad = torch.zeros((V, D), dtype=dtype, device=device) + if use_main_grad: + weights.main_grad = torch.zeros((V, D), dtype=dtype, device=device) + grad = weights.main_grad + else: + grad = weights.grad + + 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) + grad_ref = weights.grad.clone().detach() + + grad.zero_() + y = triton_embedding_forward(input_ids, weights.data_ptr(), D, dtype) + output_check(y_ref, y, name='y') + + triton_embedding_backward(dy, input_ids, grad.data_ptr(), grad.dtype) + output_check(grad_ref, grad.to(dtype), name='grad') + + grad.zero_() + 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, grad.to(dtype), name='grad') + + if bench: + ref_bytes = B * M * D * 4 + ref_time = benchmark_func(embedding.forward, input_ids) + benchmark_func(fused_accumulation_embedding_lookup, input_ids, weights, + grad_name=grad_name, + ref_time=ref_time, ref_bytes=ref_bytes) + + ref_time = benchmark_func(y_ref.backward, dy, retain_graph=True) + benchmark_func(y.backward, dy, retain_graph=True, + ref_time=ref_time, ref_bytes=ref_bytes) + + +if __name__ == '__main__': + test_scan(M=8192, bench=False) + test_embedding(B=1, M=8192, V=150000, D=8192, transpose=False, bench=False) + test_embedding(B=2, M=4096, V=150000, D=4096, transpose=True, bench=False) + test_fused_embedding(B=1, M=8192, V=150000, D=8192, transpose=False, + bench=False) + test_fused_embedding(B=1, M=4096, V=150000, D=8192, transpose=False, + bench=False) + test_fused_embedding(B=2, M=4096, V=150000, D=8192, transpose=True, + bench=False) + test_fused_embedding(B=0, M=4096, V=150000, D=8192, transpose=True, + bench=False) diff --git a/tests/test_fp32_gemm.py b/tests/test_fp32_gemm.py index 3f60d18..20f9714 100644 --- a/tests/test_fp32_gemm.py +++ b/tests/test_fp32_gemm.py @@ -5,91 +5,147 @@ import torch +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_scaled_fp32_gemm, - triton_scaled_fp32_gemm_for_update) + 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.tools.benchmark import benchmark_func -from linghe.tools.util import output_check +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 dtype = torch.bfloat16 device = 'cuda:0' - n_repeat = 100 - 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, requires_grad=True) 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) + + 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) + + y = triton_split_fp32_gemm(x, w) + 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-2) + + x.grad = None + w.grad = None + y = fp32_gemm(x, w) + y.backward(gradient=dy) + dx = x.grad + dw = w.grad + 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) - 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') + if bench: + ref_bytes = M * K * 6 + N * K * 6 + M * N * 4 + ref_flops = 2 * M * N * K + ref_time = benchmark_func(torch_fp32_matmul, x, w, + ref_bytes=ref_bytes, + ref_flops=ref_flops) + benchmark_func(triton_fp32_gemm, x, w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, ref_time=ref_time) + benchmark_func(triton_split_fp32_gemm, x, w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, ref_time=ref_time) + + ref_bytes = M * K * 10 + N * K * 4 + M * N * 4 + ref_time = benchmark_func(torch_fp32_matmul_backward, dy, w.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops) + benchmark_func(triton_fp32_gemm_for_backward, dy, w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, ref_time=ref_time) + benchmark_func(triton_split_fp32_gemm_for_backward, dy, w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, ref_time=ref_time) - 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') + ref_bytes = M * K * 4 + N * K * 12 + M * N * 4 + ref_time = benchmark_func(torch_fp32_matmul_update, dy, x.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops) + benchmark_func(triton_fp32_gemm_for_update, dy, x, + ref_bytes=ref_bytes, + ref_flops=ref_flops, ref_time=ref_time) + benchmark_func(triton_split_fp32_gemm_for_update, dy, x, + ref_bytes=ref_bytes, + ref_flops=ref_flops, ref_time=ref_time) - 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') +def test_BMK_fp32_matmul(B=2, M=2048, N=256, K=8192, bench=False): + # 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) + 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=1e-1, rtol=2e-2) + output_check(dw_ref, dw, name='update', atol=1e-1, rtol=2e-2) 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_flops=2 * M * N * K) + benchmark_func(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) + ref_flops=2 * M * N * K, ref_time=ref_time) if __name__ == '__main__': - test_fp32_matmul(M=2048, N=256, K=8192) + test_fp32_matmul(M=4096, N=256, K=8192, bench=False) + test_fp32_matmul(M=16384, N=256, K=2048, bench=False) + test_fp32_matmul(M=128, N=16, K=128, bench=False) + test_BMK_fp32_matmul(B=2, M=2048, N=16, K=8192, bench=False) + test_BMK_fp32_matmul(B=2, M=2048, N=256, K=8192, bench=False) + test_BMK_fp32_matmul(B=2, M=128, N=16, K=128, bench=False) diff --git a/tests/test_gate.py b/tests/test_gate.py new file mode 100644 index 0000000..34fcb5a --- /dev/null +++ b/tests/test_gate.py @@ -0,0 +1,140 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import torch +import torch.nn.functional as F + +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check +from linghe.utils.gate import (triton_group_rms_norm_gate_forward, + triton_group_rms_norm_gate_backward) + + +# @torch.compile +def torch_group_rms_norm_gate_forward(x, gate, weight, eps=1e-6, group_size=4, + transpose=True): + 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) + 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) + if transpose: + outputs = outputs.transpose(0, 1) + gate = F.sigmoid(gate) + outputs = (outputs * gate).to(dtype) + return outputs + + +def torch_group_rms_norm_gate_backward(grad_output, x, gate, weight, eps=1e-6, + group_size=4, + transpose=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, + transpose=transpose) + y.backward(gradient=grad_output) + return x.grad.to(dtype), gate.grad.to(dtype), weight.grad.to(dtype) + + +def test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, + transpose=True, share=False, coef=1.0, + grad_coef=1.0, + bench=False): + dtype = torch.bfloat16 + device = 'cuda:0' + x = torch.randn(bs, length, 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 transpose: + gate = (torch.randn(length, bs, dim, dtype=dtype, + device=device) * coef).requires_grad_() + grad_output = torch.randn(length, bs, dim, dtype=dtype, + device=device) * grad_coef + else: + gate = (torch.randn(bs, length, dim, dtype=dtype, + device=device) * coef).requires_grad_() + grad_output = torch.randn(bs, length, dim, dtype=dtype, + device=device) * grad_coef + + 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, name='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, 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') + + if bench: + benchmark_func(torch_group_rms_norm_gate_forward, x, gate, weight, + group_size=group_size, transpose=transpose, + ref_bytes=bs * length * dim * 6) + + benchmark_func(triton_group_rms_norm_gate_forward, x, gate, weight, + group_size=group_size, transpose=transpose, + ref_bytes=bs * length * dim * 6) + + benchmark_func(triton_group_rms_norm_gate_backward, grad_output, x, + gate, + weight, group_size=group_size, transpose=transpose, + ref_bytes=bs * length * dim * 10) + + +if __name__ == '__main__': + 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=2, length=4096, dim=2048, group_size=4, + transpose=False, share=True, + bench=False) + test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, + bench=False) + test_group_rms_norm_gate(bs=2, length=4096, dim=1536, group_size=4, + transpose=True, + bench=False) + test_group_rms_norm_gate(bs=2, length=4096, dim=1536, group_size=4, + transpose=False, + coef=10000.0, + grad_coef=10000.0, + bench=False) + test_group_rms_norm_gate(bs=2, length=4096, dim=1536, group_size=4, + transpose=False, + coef=0.0, + grad_coef=0.0, + bench=False) diff --git a/tests/test_gather.py b/tests/test_gather.py index 4606284..c4f7e4e 100644 --- a/tests/test_gather.py +++ b/tests/test_gather.py @@ -5,48 +5,59 @@ 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.quant.block import triton_batch_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 +from linghe.tools.util import (torch_batch_smooth_quant, + torch_blockwise_quant, + torch_make_indices, + torch_smooth_quant, + torch_mxfp8_quant) +from linghe.utils.gather import (triton_make_row_id_map, + triton_make_row_id_map_and_index, + 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, + triton_batch_block_pad_permute_with_indices, + triton_batch_mxfp8_permute_with_indices + ) 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) + 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) @@ -73,9 +84,11 @@ def torch_smooth_permute_with_indices(grad_data, grad_scale, indices, 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] + 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] scale = y_smooth.abs().amax(1) / 448 @@ -91,12 +104,13 @@ def torch_smooth_permute_with_indices(grad_data, grad_scale, indices, 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): +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): M, DIM = x_q.shape q_refs = [] scale_refs = [] @@ -104,21 +118,24 @@ def torch_batch_transpose_smooth_permute_with_indices(x_q, x_scale, org_smooth_s 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_q.device) + 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 + 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]] 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, y_max = 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,6 +144,103 @@ def torch_batch_transpose_smooth_permute_with_indices(x_q, x_scale, org_smooth_s return q_ref, scale_ref +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] + + 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) + 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 test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): dtype = torch.bfloat16 device = 'cuda:0' @@ -138,20 +252,18 @@ def test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): 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): + 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) @@ -214,15 +326,19 @@ def test_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8, 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) + 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) + 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, @@ -238,13 +354,15 @@ def test_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8, 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) + 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, + 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, + tokens_per_expert=token_count_per_expert, n_repeat=n_repeat, ref_time=ref_time, ref_bytes=ref_bytes) @@ -265,7 +383,7 @@ def test_triton_smooth_permute_with_mask_map(M=4096, N=4096, n_experts=32, out_tokens = sum(token_count_per_expert_list) B = 128 - grad_data = torch.randn((M, N), dtype=torch.bfloat16, device=device).to( + grad_data = torch.randn((M, N), dtype=dtype, 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, @@ -273,24 +391,23 @@ def test_triton_smooth_permute_with_mask_map(M=4096, N=4096, n_experts=32, 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') - - + 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, y_q, name='data', rtol=0.125) + output_check(scale_ref, y_scale, name='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(q_ref.float(), permuted_data.float(), name='smoothed.data', + rtol=0.125) output_check(scale_ref.float(), permuted_scale.float(), 'smoothed.scale') q_ref, scale_ref = torch_smooth_permute_with_indices(grad_data, None, @@ -300,11 +417,10 @@ def test_triton_smooth_permute_with_mask_map(M=4096, N=4096, n_experts=32, 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(q_ref.float(), permuted_data.float(), name='smoothed.data', + rtol=0.125) 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, @@ -316,16 +432,16 @@ def test_triton_smooth_permute_with_mask_map(M=4096, N=4096, n_experts=32, 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): - +def test_triton_batch_transpose_smooth_permute_with_indices(M=1024, N=2048, + n_experts=32, + topk=8, + bench=False): device = 'cuda:0' if True: logits = torch.randn((M, n_experts), dtype=torch.float32, - device=device) ** 3 - logits[:,0] -= 1000 - logits[:,2] -= 100 + 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) @@ -335,8 +451,10 @@ def test_triton_batch_transpose_smooth_permute_with_indices(M=1024, N=2048, n_ex x = torch.randn((M, N), dtype=torch.bfloat16, device=device).to( 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 + 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 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') @@ -349,56 +467,196 @@ def test_triton_batch_transpose_smooth_permute_with_indices(M=1024, N=2048, n_ex 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_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(), name='smoothed.data', rtol=0.125) + 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, scale, org_smooth_scale, smooth_scales, - indices, - token_count_per_expert_list, - round_scale=True) + 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') - 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') + 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) - 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) +def test_batch_block_pad_permute_with_indices(M=16384, N=2048, n_experts=32, + topk=2, bench=False): + 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, 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') + 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') + + if bench: + benchmark_func(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_func(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_func(triton_batch_blockwise_quant, xs, token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=num_out_tokens * N * 4) + + +def test_batch_mxfp8_permute_with_indices(M=16384, N=2048, n_experts=32, topk=2, + bench=False): + 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_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, + indices, + token_count_per_expert_list, + probs=probs) + 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') 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) + num_out_tokens = sum(token_count_per_expert_list) + benchmark_func(triton_batch_mxfp8_permute_with_indices, x, + token_count_per_expert, + indices, + token_count_per_expert_list, + probs=probs, + ref_bytes=num_out_tokens * N * 4) + + benchmark_func(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_func(triton_batch_mxfp8_quant, xs, token_count_per_expert, + token_count_per_expert_list, + ref_bytes=num_out_tokens * N * 4) 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_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) + 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) + + test_batch_block_pad_permute_with_indices(M=8192 * 2, N=2048, n_experts=32, + topk=2, bench=False) + test_batch_block_pad_permute_with_indices(M=0, N=2048, n_experts=32, topk=2, + bench=False) + test_batch_block_pad_permute_with_indices(M=8192, N=1536, n_experts=32, + topk=2, bench=False) + + test_batch_mxfp8_permute_with_indices(M=3095, N=2048, n_experts=32, topk=2, + bench=False) + test_batch_mxfp8_permute_with_indices(M=0, N=2048, n_experts=32, topk=2, + bench=False) + test_batch_mxfp8_permute_with_indices(M=1024, N=1536, n_experts=32, topk=2, + bench=False) diff --git a/tests/test_group_quant.py b/tests/test_group_quant.py index d8c68ec..b7c4c26 100644 --- a/tests/test_group_quant.py +++ b/tests/test_group_quant.py @@ -7,23 +7,25 @@ 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 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) + 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) + test_group_quant(M=2049, N=1536, B=128) diff --git a/tests/test_hadamard_quant.py b/tests/test_hadamard_quant.py index f508e2f..ae9263b 100644 --- a/tests/test_hadamard_quant.py +++ b/tests/test_hadamard_quant.py @@ -5,29 +5,30 @@ 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): + 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) + 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) + 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 + dtype = torch.bfloat16 device = 'cuda:0' x = torch.randn((M, K), dtype=dtype, device=device) w = torch.randn((N, K), dtype=dtype, device=device) @@ -35,47 +36,44 @@ def test_hadamard_quant(M=8192, N=1024, K=2048, B=64, bench=False): 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(xq, x_q, 'x.data', rtol=0.125) output_check(xs, x_scale, 'x.scale') - output_check(xqt, xt_q, 'xt.data') + 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(wq, w_q, 'w.data', rtol=0.125) output_check(ws, w_scale, 'w.scale') - output_check(wqt, wt_q, 'wt.data') + 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(dyq, dy_q, 'dy.data', rtol=0.125) output_check(dys, dy_scale, 'dy.scale') - output_check(dyqt, dyt_q, 'dyt.data') + 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 + 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_() @@ -83,18 +81,19 @@ def test_hadamard_quant_linear(M=8192, N=1024, K=2048, B=64): 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__': 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..76e7a9f --- /dev/null +++ b/tests/test_la.py @@ -0,0 +1,200 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math + +import torch + +from linghe.attn.la import (triton_lightning_attention_forward, + triton_lightning_attention_backward, + triton_fused_lightning_attention_backward) +from linghe.tools.benchmark import benchmark_func +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 + + +def test_la(bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, + bench=False): + 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) + + if bench: + ref_bytes = bs * length * qo_heads * dim * 8 + bs * qo_heads * dim * dim * 8 + benchmark_func(triton_lightning_attention_forward, q, k, v, + decay_scales, ref_bytes=ref_bytes) + benchmark_func(triton_lightning_attention_backward, g, q, k, v, + decay_scales, ref_bytes=ref_bytes) + benchmark_func(triton_fused_lightning_attention_backward, g, q, k, v, + state, decay_scales, ref_bytes=ref_bytes) + + +# def test_varlen_la(qls=[1024,1024], qo_heads=16, kv_heads=16, dim=128, digest=False, bench=False): +# device = torch.device('cuda:0') +# dtype = torch.bfloat16 +# kls = qls + +# assert all([x<=kls[i] for i,x in enumerate(qls)]) +# bs = len(qls) + +# ref_bytes = sum(qls) * qo_heads * dim * 8 + bs * qo_heads * dim * dim * 8 + +# q, k, v = make_input(qo_heads=qo_heads, kv_heads=kv_heads, dim=dim, qls=qls, kls=kls) + +# 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 = 2**(-0.5 * torch.ones(qo_head, dtype=torch.float32, device=device)) +# lengths = torch.tensor([0] + qls, device=device, dtype=torch.long) +# cu_seqlens = torch.cumsum(lengths, 0) +# padded_cu_seqlens = cu_seqlens + +# output_ref, state_ref = torch_varlen_linear_attn(q, k, v, s, decay_scales, cu_seqlens, padded_cu_seqlens) + +# max_q_length = max(qls) +# output = triton_lightning_attention_forward(q, k, v, decay_scales, cu_seqlens, padded_cu_seqlens, max_q_length) + +# output_check(output_ref, output, name='output', rtol=0.1, atol=0.1) +# output_check(state_ref, s, name='state', rtol=0.01, atol=0.01) + +# if digest: +# print( +# f"output_ref max:{torch.max(output_ref).item():.3f} min:{torch.min(output_ref).item():.3f}") +# print( +# f"output max:{torch.max(output).item():.3f} min:{torch.min(output).item():.3f}") + +# print("output_ref[:,0,0]", output_ref[:, 0, 0]) +# print("output[:,0,0]", output[:, 0, 0]) + +# print("output_ref[0,:,0]", output_ref[0, :, 0]) +# print("output[0,:,0]", output[0, :, 0]) + +# print("output_ref[0,0,:]", output_ref[0, 0, :]) +# print("output[0,0,:]", output[0, 0, :]) + +# if bench: +# benchmark_func(triton_lightning_attention_forward, q, k, v, decay_scales, cu_seqlens, padded_cu_seqlens, max_q_length, ref_bytes=ref_bytes) + + +if __name__ == '__main__': + test_la(bs=1, length=8192, qo_heads=64, kv_heads=64, dim=128, digest=False, + bench=True) diff --git a/tests/test_loss.py b/tests/test_loss.py index 0ba2631..256f61b 100644 --- a/tests/test_loss.py +++ b/tests/test_loss.py @@ -7,58 +7,163 @@ import torch +from linghe.facade.loss import moe_z_loss, 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_forward, \ - triton_softmax_cross_entropy_backward +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 -def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, bench=False): +def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, grad_coef=1.0, + ignore_index=None, fill=False, + inplace=False, bench=False): device = 'cuda:0' - logits = torch.randn((M, N), dtype=torch.bfloat16, device=device, + 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) + + 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, + 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, - input_grad, output_grad=logits) - output_check(grad_ref.float(), grad.float(), mode='grad') + 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) + if bench: - benchmark_func(torch_cross_entropy, logits, targets, - ref_bytes=M * N * 6) + benchmark_func(torch_cross_entropy, logits.requires_grad_(), targets, + output_grad, + ref_bytes=M * N * 2) benchmark_func(triton_softmax_cross_entropy_forward, logits, targets, + ignore_index=ignore_index, 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) + benchmark_func(triton_softmax_cross_entropy_backward, + logits.detach().clone(), targets, + sum_exp, max_logit, output_grad, + ignore_index=ignore_index, ref_bytes=M * N * 4) + + +def test_z_loss(L=4096, B=2, N=256, coef=0.001, bench=False): + 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') + + if bench: + benchmark_func(torch_z_loss, logits, coef=coef, + ref_bytes=L * B * N * 4) + benchmark_func(triton_moe_z_loss_forward, logits, coef=coef, + ref_bytes=L * B * N * 4) + benchmark_func(triton_moe_z_loss_backward, input_grad, logits, + coef=coef, ref_bytes=L * B * N * 8) 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) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, + inplace=True, bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, + grad_coef=1e-6, inplace=True, bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=10000.0, + grad_coef=100.0, fill=True, inplace=True, + bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, + fill=True, ignore_index=-100, + inplace=True, bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, + fill=True, ignore_index=0, inplace=True, + bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184 - 16, coef=10000.0, + grad_coef=100.0, fill=True, inplace=True, + bench=True) + test_triton_softmax_cross_entropy(M=8192, N=175175, coef=1.0, grad_coef=1.0, + inplace=True, bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=0.0, grad_coef=0.0, + inplace=True, bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=0.0, + grad_coef=100.0, inplace=True, bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1000.0, + grad_coef=0.0, inplace=True, bench=True) + test_triton_softmax_cross_entropy(M=8192, N=157184, coef=100.0, + grad_coef=100.0, fill=True, inplace=True, + bench=True) + test_triton_softmax_cross_entropy(M=4096, N=157184, coef=0.1, grad_coef=1.0, + inplace=True, bench=False) + + test_z_loss(L=4096, B=2, N=256, coef=1e-6, bench=False) diff --git a/tests/test_mla.py b/tests/test_mla.py new file mode 100644 index 0000000..a213863 --- /dev/null +++ b/tests/test_mla.py @@ -0,0 +1,352 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import math + +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.benchmark import benchmark_func +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 + + +def test_softmax(M=128, N=128): + 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') + + +def test_dot_sum(M=128, N=128, D=128): + 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') + + +def test_mla(B=2, L=4096, H=16, causal=True, hpc=False, safe=True, coef=1.0, + clip_value=None, bench=False): + 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') + + if bench: + ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) + benchmark_func(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_func(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) + + +def test_varlen_mla(LS=[2048, 4096], H=16, causal=True, hpc=False, safe=True, + coef=1.0, clip_value=None, pad=False, bench=False): + 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=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') + + if bench: + ref_flops = sum( + [L * L * H * (192 + 128) * (1 if causal else 2) for L in LS]) + benchmark_func(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_func(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) + + +def test_fp8_mla(B=2, L=4096, H=16, causal=True, hpc=False, quant_value=False, + bench=False): + 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') + + if bench: + ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) + benchmark_func(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) + + +if __name__ == "__main__": + # test_softmax(M=128, N=128) + + # test_dot_sum(M=128, N=128, D=128) + + test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=False, coef=1.0, + clip_value=500, bench=True) + # test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=500.0, bench=True) + # test_mla(B=1, L=8192, H=64, causal=True, hpc=True, safe=False, coef=1.0, clip_value=None, bench=False) + # test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=True, coef=100.0, clip_value=None, bench=False) + # test_mla(B=1, L=4096, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + # test_mla(B=1, L=4096, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + # test_mla(B=1, L=8192, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + # test_mla(B=1, L=8192, H=1, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + + # test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, pad=False, bench=True) + # test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=100.0, pad=False, bench=True) + # test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=None, pad=True, bench=True) + + # test_varlen_mla(LS=[4096,4096], H=64, causal=True, hpc=False, safe=True, coef=1.0, bench=False) + # test_varlen_mla(LS=[2048,2048,4096], H=64, causal=True, hpc=True, safe=True, coef=1.0, bench=False) + # test_varlen_mla(LS=[127,873,3096], H=64, causal=False, hpc=False, safe=False, coef=1.0, bench=False) + # test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, clip_value=100.0, bench=False) + # test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, pad=True, bench=False) + # test_varlen_mla(LS=[127,873,3456], H=1, causal=True, hpc=False, safe=True, coef=1.0, bench=False) + + # test_fp8_mla(B=1, L=8192, H=64, causal=True, hpc=False, quant_value=False, bench=True) diff --git a/tests/test_mul.py b/tests/test_mul.py new file mode 100644 index 0000000..104d190 --- /dev/null +++ b/tests/test_mul.py @@ -0,0 +1,104 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import random + +import torch + +from linghe.tools.benchmark import benchmark_func +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] + + +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_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) + + if bench: + ref_time = benchmark_func(torch_fp16_dot, x, y, n_repeat=n_repeat) + ref_time = benchmark_func(triton_dot, x, q, n_repeat=n_repeat, + ref_time=ref_time) + + +def test_inplace_scale(M=2 ** 20, bench=False): + 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 + + if bench: + ref_time = benchmark_func(torch_inplace_scale, x, scale, + ref_bytes=ref_bytes) + benchmark_func(triton_inplace_scale, x, scale, + ref_bytes=ref_bytes, ref_time=ref_time) + + +def test_batch_scale(M=4096, N=2048, k=128, scale=1.0, bench=False): + 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)] + # xs.append(torch.randn(2**32//N, N, + # dtype=dtype, device='cuda:0')) + 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 + + if bench: + ref_time = benchmark_func(torch_batch_scale, xs, scale, + ref_bytes=ref_bytes) + benchmark_func(triton_batch_scale, xs, scale, + ref_bytes=ref_bytes, ref_time=ref_time) + + +if __name__ == '__main__': + test_dot(M=4096, N=4096, bench=False) + test_inplace_scale(M=2 ** 28 + 1, bench=False) + test_batch_scale(M=2048, N=1024, k=128, scale=2.0, bench=False) + test_batch_scale(M=2048, N=1024, k=128, scale=0.0, bench=False) diff --git a/tests/test_mxfp8_quant.py b/tests/test_mxfp8_quant.py new file mode 100644 index 0000000..2054306 --- /dev/null +++ b/tests/test_mxfp8_quant.py @@ -0,0 +1,98 @@ +# -*- coding: utf-8 -*- +""" +Copyright (c) Ant Financial Service Group and its affiliates. +""" + +import random + +import torch + +from linghe.quant.mxfp8 import triton_mxfp8_quant, triton_batch_mxfp8_quant +from linghe.tools.benchmark import benchmark_func +from linghe.tools.check import output_check +from linghe.tools.util import torch_mxfp8_quant + + +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 test_mxfp8_quant(M=4096, N=4096, bench=False): + dtype = torch.bfloat16 + device = 'cuda:0' + + x = torch.randn(M, N, dtype=dtype, device=device) + + x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = torch_mxfp8_quant(x) + x_q, x_scale, xt_q, xt_scale = triton_mxfp8_quant(x) + + 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 * 4 + benchmark_func(triton_mxfp8_quant, x, ref_bytes=ref_bytes) + benchmark_func(torch_mxfp8_quant, x) + + +def test_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] + token_count_per_expert = torch.tensor(splits, device=device) + + x = torch.randn((sum(splits), N), dtype=dtype, device=device) + + x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = torch_batch_mxfp8_quant(x, + splits) + + 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(torch_batch_mxfp8_quant, x, splits) + benchmark_func(triton_batch_mxfp8_quant, x, token_count_per_expert, + splits, output_mode=2, ref_bytes=ref_bytes) + + +if __name__ == '__main__': + test_mxfp8_quant(M=4096, N=8192, bench=False) + test_mxfp8_quant(M=4031, N=8192, bench=False) + test_mxfp8_quant(M=4096, N=8192, bench=False) + test_batch_mxfp8_quant(M=4096, N=8192, bench=False) diff --git a/tests/test_norm.py b/tests/test_norm.py index c94f455..61bb916 100644 --- a/tests/test_norm.py +++ b/tests/test_norm.py @@ -4,21 +4,24 @@ """ import torch -import torch.nn.functional as F +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, + torch_mxfp8_quant) from linghe.utils.norm import (triton_rms_norm_and_smooth_quant_forward, triton_rms_norm_and_block_quant_forward, + triton_rms_norm_and_mxfp8_quant_forward, + triton_rms_norm_fp32_gemm_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.tools.benchmark import benchmark_func + triton_rms_norm_forward) def torch_rms_forward(x, weight): + dtype = x.dtype x = x.float() weight = weight.float() N = x.shape[-1] @@ -30,10 +33,11 @@ def torch_rms_forward(x, weight): ) 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() @@ -49,7 +53,7 @@ def torch_rms_backward(x, weight, dy): 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, @@ -86,48 +90,61 @@ def torch_rms_and_block_quant_forward(x, weight, round_scale=False): 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() + 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 torch_rms_and_mxfp8_quant_forward(x, weight): 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 + 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_(weight) + y = rmsnorm(x) + # mxfp8 + y_q, y_scale, yt_q, yt_scale = torch_mxfp8_quant(y) + return y_q, y_scale, yt_q, yt_scale def test_rmsnorm(M=4096, N=4096, bench=False): @@ -139,18 +156,29 @@ def test_rmsnorm(M=4096, N=4096, bench=False): 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_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): @@ -173,10 +201,10 @@ def test_rmsnorm_and_smooth_quant(M=4096, N=4096, bench=False): calibrate=calibrate, output_rms=True, round_scale=True) - output_check(q_ref, q, mode="smooth.data") - output_check(scale_ref, scale, mode='smooth.scale') + 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, @@ -195,30 +223,31 @@ def test_rmsnorm_and_block_quant(M=4096, N=4096, bench=False): 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, + q_ref, scale_ref, rms_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, 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, mode='0.block.t_data') - output_check(scale_t_ref.t(), scale_t, mode="0.block.t_scale") + 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") + + 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(triton_rms_norm_and_block_quant_forward, x, weight, @@ -229,82 +258,134 @@ def test_rmsnorm_and_block_quant(M=4096, N=4096, bench=False): 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_rmsnorm_and_mxfp8_quant(M=4096, N=4096, bench=False): + dtype = torch.bfloat16 + 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) + + # mxfp8 + q_ref, scale_ref, qt_ref, scale_t_ref = torch_rms_and_mxfp8_quant_forward(x, + weight) + q, scale, rms, q_t, scale_t = triton_rms_norm_and_mxfp8_quant_forward(x, + weight, + 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") + + q, scale, _, _, _ = triton_rms_norm_and_mxfp8_quant_forward(x, weight, + 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_mxfp8_quant_forward(x, weight, + rms=rms, + output_mode=1) + output_check(qt_ref, q_t, name='0.block.t_data', rtol=0.125) + output_check(scale_t_ref, scale_t, name="0.block.t_scale") + + if bench: + benchmark_func(triton_rms_norm_and_mxfp8_quant_forward, x, weight, + output_mode=0, + ref_bytes=M * N * 3) + + benchmark_func(triton_rms_norm_and_mxfp8_quant_forward, x, weight, + output_mode=1, + ref_bytes=M * N * 3) + + benchmark_func(triton_rms_norm_and_mxfp8_quant_forward, x, weight, output_mode=2, ref_bytes=M * N * 4) -def test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, - transpose=True, - bench=False): +def test_rms_norm_fp32_gemm_block_quant_forward(M=8192, N=256, K=2048, + 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') + + 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 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_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_group_rms_norm_gate_forward, x, gate, weight, - group_size=group_size, - ref_bytes=bs * length * dim * 6) + 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(triton_group_rms_norm_gate_backward, grad_output, x, gate, - weight, group_size=group_size, - ref_bytes=bs * length * dim * 10) + 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=True) + + test_rmsnorm_and_mxfp8_quant(M=2048, N=1664, bench=False) + test_rmsnorm_and_mxfp8_quant(M=8192, N=4096, 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..2490daf 100644 --- a/tests/test_rearange.py +++ b/tests/test_rearange.py @@ -6,11 +6,11 @@ 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,16 +22,7 @@ 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 - +def test_sort_chunks_by_index(M=4096, N=4096, bench=False): dtype = torch.bfloat16 device = 'cuda:0' n_repeat = 100 @@ -49,15 +40,17 @@ 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_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_split_and_cat(x_q, counts, indices, scales=x_scales) + data, scale = triton_sort_chunks_by_index(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') + output_check(data_ref.view(torch.float8_e4m3fn), data, + name='data') + output_check(scale_ref, scale, name='scale') if bench: benchmark_func(torch.split, x_q.view(torch.uint8), split_size_list, @@ -66,12 +59,13 @@ def test_triton_split_and_cat(M=4096, N=4096, bench=False): 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), + benchmark_func(torch_sort_chunks_by_index, + 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, + benchmark_func(triton_sort_chunks_by_index, x_q, counts, indices, scales=x_scales, n_repeat=n_repeat) if __name__ == '__main__': - test_triton_split_and_cat(M=4096, N=4096) + test_sort_chunks_by_index(M=4096, N=4096) diff --git a/tests/test_reduce.py b/tests/test_reduce.py index bff2b65..59617e8 100644 --- a/tests/test_reduce.py +++ b/tests/test_reduce.py @@ -3,17 +3,28 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import random + import torch from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check +from linghe.tools.check import output_check from linghe.utils.reduce import (triton_abs_max, - triton_batch_count_zero, - triton_batch_sum_with_ord) + 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): @@ -47,10 +58,6 @@ 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, @@ -59,25 +66,57 @@ def test_count_zero(M=4096, N=8192, k=32, bench=False): 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)] +def test_norm(M=4096, N=8192, coef=1.0, bench=False): + x = torch.randn(M, N, dtype=torch.float32, device='cuda:0') * 1.0 - ref_bytes = sum([x.numel() for x in xs]) * 4 + 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') + + if bench: + ref_bytes = M * N * 4 + n_repeat = 100 + ref_time = benchmark_func(lambda x: x.norm(p=2), x, n_repeat=n_repeat, + ref_bytes=ref_bytes) + benchmark_func(triton_norm, x, ord=2, norm=True, scalar=True, + n_repeat=n_repeat, + ref_bytes=ref_bytes, ref_time=ref_time) + + +def test_batch_norm(M=4096, N=8192, k=32, coef=1.0, bench=False): + 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) - sums = triton_batch_sum_with_ord(xs) - output_check(sum_ref, sums, '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') if bench: + ref_bytes = sum([x.numel() for x in xs]) * 4 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, + benchmark_func(triton_batch_norm, 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) + test_triton_abs_max(M=4096, N=4096, bench=False) + test_count_zero(M=4096, N=8192, k=32, bench=False) + test_norm(M=100000, N=8192, bench=False) + test_batch_norm(M=4096, N=1024, k=16, bench=False) + test_batch_norm(M=4096, N=1024, k=64, bench=False) + test_batch_norm(M=4096, N=2048, k=1024, coef=1e12, bench=True) diff --git a/tests/test_rope.py b/tests/test_rope.py index f3ffd31..61b0b55 100644 --- a/tests/test_rope.py +++ b/tests/test_rope.py @@ -5,11 +5,17 @@ import torch +from linghe.facade.rope import qk_norm_half_rope 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.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): @@ -20,8 +26,14 @@ def rotate_half(x): 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 @@ -29,20 +41,21 @@ def apply_rotary_pos_emb(q, k, cos, sin, position_ids): 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)) + 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) else: @@ -51,7 +64,7 @@ def torch_half_rope(q, k, freqs, rope_theta=10000.0, transposed=True): 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 +82,128 @@ 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,42 +222,87 @@ 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 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 def test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, - transposed=True, + transposed=True, bench=False): - dtype = torch.bfloat16 + 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, @@ -140,9 +311,11 @@ def test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, def test_qk_norm_and_half_rope(B=2, L=4096, H=32, h=8, D=128, - rope_theta=10000.0, + rope_theta=10000.0, + eps=1e-6, interleaved=True, transposed=True, + silu=False, bench=False): dtype = torch.bfloat16 device = 'cuda:0' @@ -150,78 +323,392 @@ def test_qk_norm_and_half_rope(B=2, L=4096, H=32, h=8, D=128, 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, + 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) + 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=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, + h=h, eps=eps, 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 + 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=1e-6, + freqs, eps=eps, 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') + 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') 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, + 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_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, + 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) +def test_varlen_qk_norm_and_half_rope(lengths=[2048, 2048], H=32, h=4, dim=128, + rope_theta=10000.0, silu=False, + interleaved=True, + bench=False, cp_size=1, cp_rank=0): + dtype = torch.bfloat16 + 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=5.0, rtol=0.02) + output_check(dkw_ref, dkw.to(dtype), name='dkw', atol=5.0, rtol=0.02) + + if bench: + lbh = sum(lengths) // cp_size * H + benchmark_func(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_func(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) + + +def test_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=False, + bench=False): + 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') + + if bench: + lbh = L * B * H + benchmark_func(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_func(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) + + +def test_varlen_mla_rope(lengths=[2048, 2048], H=32, rope_theta=10000.0, + bench=False, cp_size=1, cp_rank=0): + dtype = torch.bfloat16 + 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) + + if bench: + lbh = sum(lengths) // cp_size * H + benchmark_func(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_func(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) + + if __name__ == '__main__': - test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, transposed=True, + 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, + 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=2, L=4096, H=16, h=16, D=128, + rope_theta=10000.0, interleaved=True, + transposed=True, silu=True, bench=False) + test_qk_norm_and_half_rope(B=2, L=4096, H=16, h=16, D=128, + rope_theta=10000.0, interleaved=True, + transposed=True, silu=False, 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=False, silu=True, bench=True) 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) + rope_theta=10000.0, interleaved=True, + transposed=False, silu=False, bench=False) + test_qk_norm_and_half_rope(B=4, L=4096, H=32, h=4, D=128, + rope_theta=10000.0, interleaved=False, + transposed=True, silu=True, bench=False) + test_qk_norm_and_half_rope(B=4, L=4096, H=24, h=6, D=128, + rope_theta=10000.0, interleaved=True, + transposed=True, silu=False, bench=False) + test_qk_norm_and_half_rope(B=4, L=4096, H=32, h=32, D=128, + rope_theta=10000.0, interleaved=False, + transposed=False, silu=True, bench=False) + test_qk_norm_and_half_rope(B=1, L=4096, H=32, h=32, D=128, + rope_theta=10000.0, interleaved=False, + transposed=False, silu=False, bench=False) + test_varlen_qk_norm_and_half_rope(lengths=[2048], H=24, h=6, dim=128, + rope_theta=10000.0, silu=False, + interleaved=True, cp_size=1, cp_rank=0, + bench=False) + test_varlen_qk_norm_and_half_rope(lengths=[1024, 4096, 4096, 568], H=32, + h=4, dim=128, rope_theta=10000.0, + silu=False, + interleaved=True, cp_size=1, cp_rank=0, + bench=False) + test_varlen_qk_norm_and_half_rope(lengths=[2048, 4096, 4096], H=32, h=4, + dim=128, rope_theta=10000.0, silu=False, + interleaved=True, cp_size=1, cp_rank=0, + bench=False) + test_varlen_qk_norm_and_half_rope(lengths=[2048, 3072, 4096], H=32, h=4, + dim=128, rope_theta=10000.0, silu=False, + interleaved=True, cp_size=4, cp_rank=0, + bench=False) + test_varlen_qk_norm_and_half_rope(lengths=[2048, 4096, 4096], H=32, h=4, + dim=128, rope_theta=10000.0, silu=True, + interleaved=False, cp_size=4, cp_rank=0, + bench=False) + test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=False, + bench=False) + test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=True, + bench=False) + test_varlen_mla_rope(lengths=[8192], H=64, rope_theta=10000.0, cp_size=1, + cp_rank=0, + bench=False) + test_varlen_mla_rope(lengths=[4096, 4096], H=16, rope_theta=10000.0, + cp_size=1, cp_rank=0, + bench=False) + test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, + rope_theta=10000.0, cp_size=4, cp_rank=0, + bench=False) + test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, + rope_theta=10000.0, cp_size=4, cp_rank=1, + bench=False) + test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, + rope_theta=10000.0, cp_size=4, cp_rank=2, + bench=False) + test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, + rope_theta=10000.0, cp_size=4, cp_rank=3, + bench=False) diff --git a/tests/test_scatter.py b/tests/test_scatter.py index fa9fbd6..89ff9b7 100644 --- a/tests/test_scatter.py +++ b/tests/test_scatter.py @@ -6,22 +6,25 @@ import torch from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import (output_check, - torch_make_indices) +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 - ) + 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 + return outputs.to(dtype) def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): @@ -39,8 +42,7 @@ 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) + sums_ref = torch_scatter_add(x, outputs.clone(), indices, None) unpermuted_prob = probs.T.contiguous().masked_select( mask_map.T.contiguous()) @@ -51,8 +53,6 @@ def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): 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, @@ -60,5 +60,6 @@ def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): 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) + test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False) + test_scatter(M=2467, N=4096, n_experts=32, topk=2, bias=-0.1, bench=False) + test_scatter(M=2467, N=1536, n_experts=32, topk=2, bias=-0.1, bench=False) diff --git a/tests/test_silu.py b/tests/test_silu.py index 2375b90..fac5c88 100644 --- a/tests/test_silu.py +++ b/tests/test_silu.py @@ -16,12 +16,19 @@ 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) -from linghe.tools.util import output_check, torch_smooth_quant, \ - torch_group_quant + 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): @@ -32,18 +39,23 @@ def torch_silu(x): def torch_weighted_silu(x, weight): + 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 + return y.to(dtype) def torch_weighted_silu_backward(dy, x, weight): + 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.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): @@ -76,6 +88,16 @@ def torch_silu_and_block_quant_forward(x, round_scale=True): return y_q, y_scale, yt_q, yt_scale +def torch_silu_and_mxfp8_quant_forward(x): + 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, 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, round_scale=True, reverse=True): @@ -107,6 +129,17 @@ def torch_silu_and_block_quant_backward(grad, x, round_scale=True): return q, dx_scale, yt_q, yt_scale +def torch_silu_and_mxfp8_quant_backward(grad, x): + grad = grad.float() + x = x.float().detach().clone().requires_grad_() + y = torch_silu(x) + 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, @@ -184,6 +217,44 @@ def torch_batch_weighted_silu_and_block_quant_forward(xs, weight, return qs, scales, qts, qtscales +def torch_batch_weighted_silu_and_mxfp8_quant_forward(xs, weight, + counts): + 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]) + + 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, @@ -249,7 +320,7 @@ 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, + qtscales = torch.zeros((0,), device=device, dtype=torch.float32) return dx_q, dx_scale, dw, qts, qtscales @@ -280,57 +351,90 @@ 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): +def torch_batch_weighted_silu_and_mxfp8_quant_backward(grad_output, x, weight, + counts): + 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) + 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 + + +def test_weighted_silu(M=4096, N=4096, asm=False, coef=1.0, 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') + x = (x * coef).clone().detach().requires_grad_() + weight = torch.randn((M, 1), dtype=torch.float32, device='cuda:0') grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, device='cuda:0') ref_y = torch_weighted_silu(x, weight) - y = triton_weighted_silu_forward(x, weight) + 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') + output_check(dw_ref, dw, 'dw', rtol=3e-3, atol=3e-3) if bench: - benchmark_func(triton_weighted_silu_forward, x, weight, n_repeat=100, + benchmark_func(triton_weighted_silu_forward, x, weight, asm=asm, + 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, +def test_silu_and_smooth_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, + bench=False): + x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') + x = (x * coef).clone().detach().requires_grad_() + grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, + device='cuda:0') * grad_coef + smooth_scale = 1 + torch.rand((N // 2,), dtype=torch.float32, 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') + 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') + round_scale = False y_q_ref, y_scale_ref, y_maxs_ref = torch_silu_and_smooth_quant_forward(x, - smooth_scale=smooth_scale) + smooth_scale=smooth_scale, + round_scale=round_scale) y_q, y_scale, y_maxs = triton_silu_and_smooth_quant_forward(x, smooth_scale=smooth_scale, - round_scale=True, + round_scale=round_scale, calibrate=True) - output_check(y_q_ref.float(), y_q.float(), 'smooth.y_q') + output_check(y_q_ref, y_q, 'smooth.y_q', rtol=0.125) output_check(y_scale_ref, y_scale, 'smooth.y_scale') output_check(y_maxs_ref, y_maxs, 'smooth.y_max') @@ -347,9 +451,9 @@ def test_silu_and_smooth_quant(M=4096, N=4096, bench=False): reverse=True, round_scale=True) - output_check(dx_q_ref.float(), dx_q.float(), 'smooth.dx_data') + 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.float(), dxt_q.float(), 'smooth.dxt_data') + output_check(dxt_q_ref, dxt_q, 'smooth.dxt_data', rtol=0.125) output_check(dxt_scale_ref, dxt_scale, 'smooth.dxt_scale') if bench: @@ -365,90 +469,127 @@ def test_silu_and_smooth_quant(M=4096, N=4096, bench=False): 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] +def test_silu_and_block_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, + bench=False): + x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') + x = (x * coef).clone().detach().requires_grad_() + grad_output = torch.randn((M, N // 2), 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') + x, round_scale=round_scale) y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward(x, - round_scale=True, + round_scale=round_scale, output_mode=0) - output_check(y_q_ref.float(), y_q.float(), 'block.0.y_q') + 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=True, + round_scale=round_scale, output_mode=1) - output_check(yt_q_ref.float(), yt_q.float(), 'block.1.yt_q') + 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, + 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=True) + round_scale=round_scale) 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') + round_scale=round_scale) + 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.float(), dxt_q.float(), 'block.dxt_q') + output_check(dxt_q_ref, dxt_q, 'block.dxt_q', rtol=0.125) output_check(dxt_scale_ref.t(), dxt_scale, 'block.dxt_scale') if bench: benchmark_func(triton_silu_and_block_quant_forward, x, + round_scale=round_scale, output_mode=0, + n_repeat=100, ref_bytes=M * N * 3) + benchmark_func(triton_silu_and_block_quant_forward, x, + round_scale=round_scale, output_mode=1, + n_repeat=100, ref_bytes=M * N * 3) + benchmark_func(triton_silu_and_block_quant_forward, x, + round_scale=round_scale, output_mode=2, 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_silu_and_mxfp8_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, + bench=False): + x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') + x = (x * coef).clone().detach().requires_grad_() + grad_output = torch.randn((M, N // 2), 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) + y_q, y_scale, yt_q, yt_scale = triton_silu_and_mxfp8_quant_forward(x, + 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, + 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, + 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) + dx_q, dx_scale, dxt_q, dxt_scale = triton_silu_and_mxfp8_quant_backward( + grad_output, x) + 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') + + if bench: + benchmark_func(triton_silu_and_mxfp8_quant_forward, x, + n_repeat=100, ref_bytes=M * N * 3) + benchmark_func(triton_silu_and_mxfp8_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, + coef=1.0, + grad_coef=1.0, 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] + 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') * coef + 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 grad_output = torch.randn((bs, N // 2), dtype=torch.bfloat16, - device='cuda:0') ** 3 + device='cuda:0') * grad_coef 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 round_scale = True - + rtol = 2 if round_scale else 0.125 x_q_ref, x_scale_ref, x_max_ref = torch_batch_weighted_silu_and_smooth_quant_forward( x, weight, @@ -462,7 +603,7 @@ def test_triton_batch_weighted_silu_and_smooth_quant(M=1024, N=4096, smooth_scale=smooth_scales, round_scale=round_scale, reverse=False) - output_check(x_q_ref.float(), x_q.float(), 'smooth.data') + 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( @@ -477,10 +618,11 @@ def test_triton_batch_weighted_silu_and_smooth_quant(M=1024, N=4096, splits=count_list, round_scale=round_scale, reverse=False) - output_check(dx_ref.float(), dx.float(), 'smooth.dx') + output_check(dx_ref, dx, 'smooth.dx', rtol=rtol) 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') + 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') if bench: @@ -500,29 +642,24 @@ def test_triton_batch_weighted_silu_and_smooth_quant(M=1024, N=4096, 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] + bench=False, + 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, 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 // 2), dtype=torch.bfloat16, - device='cuda:0') ** 3 - round_scale = True + 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, @@ -537,10 +674,30 @@ def test_triton_batch_weighted_silu_and_block_quant(M=1024, N=4096, 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') + 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') + + x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_block_quant_forward( + x, + weight, + counts, + count_list, + 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, + round_scale=round_scale, + 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, @@ -548,11 +705,12 @@ def test_triton_batch_weighted_silu_and_block_quant(M=1024, N=4096, 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') + 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') if bench: ref_time = None @@ -561,6 +719,11 @@ def test_triton_batch_weighted_silu_and_block_quant(M=1024, N=4096, 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=1, n_repeat=100, + ref_bytes=n_experts * M * N * 3, ref_time=ref_time) benchmark_func(triton_batch_weighted_silu_and_block_quant_forward, x, weight, counts, round_scale=True, splits=count_list, @@ -572,21 +735,117 @@ def test_triton_batch_weighted_silu_and_block_quant(M=1024, N=4096, ref_bytes=n_experts * M * N * 4, ref_time=ref_time) +def test_triton_batch_weighted_silu_and_mxfp8_quant(M=1024, N=4096, + n_experts=32, + coef=1.0, + grad_coef=1.0, + bench=False): + 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') * coef + weight = torch.randn((bs, 1), dtype=torch.float32, device='cuda:0') + + grad_output = torch.randn((bs, N // 2), 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) + x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_mxfp8_quant_forward( + x, + weight, + counts, + count_list, + 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) + dx, dx_scale, dw, dxt, dxt_scale = triton_batch_weighted_silu_and_mxfp8_quant_backward( + grad_output, x, weight, counts, splits=count_list) + 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) + + if bench: + ref_time = None + benchmark_func(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 * 2.5, ref_time=ref_time) + benchmark_func(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 * 3, ref_time=ref_time) + benchmark_func(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 * 4, ref_time=ref_time) + + if __name__ == '__main__': - test_weighted_silu(M=16384, N=1024, bench=True) + test_weighted_silu(M=16384, N=4096, coef=1.0, asm=False, bench=False) + test_weighted_silu(M=16384, N=4096, coef=1.0, asm=True, bench=False) + test_weighted_silu(M=8192, N=1536, bench=False) + test_weighted_silu(M=0, N=1536, bench=False) 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_silu_and_block_quant(M=16384, N=1024, bench=False) + test_silu_and_block_quant(M=8192, N=4096, bench=False) + test_silu_and_block_quant(M=16384, N=1536, bench=False) + test_silu_and_block_quant(M=4096, N=1536 * 8, bench=False) + test_silu_and_block_quant(M=4096, N=1536 * 8, coef=100.0, grad_coef=100.0, + bench=False) + test_silu_and_block_quant(M=4096, N=1536 * 8, coef=0.0, grad_coef=0.0, + bench=False) + + test_silu_and_mxfp8_quant(M=16384, N=1024, bench=False) + test_silu_and_mxfp8_quant(M=2345, N=1024, bench=False) + test_silu_and_mxfp8_quant(M=2345, N=1536, 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_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) + test_triton_batch_weighted_silu_and_block_quant(M=0, N=1536, n_experts=32, + bench=False) + test_triton_batch_weighted_silu_and_block_quant(M=2048, N=8192, + n_experts=32, bench=False) + test_triton_batch_weighted_silu_and_block_quant(M=12080, N=1536, + n_experts=32, coef=100.0, + grad_coef=100.0, + bench=False) + test_triton_batch_weighted_silu_and_block_quant(M=12080, N=1536, + n_experts=32, coef=0.0, + grad_coef=0.0, bench=False) + + test_triton_batch_weighted_silu_and_mxfp8_quant(M=0, N=2048, n_experts=32, + bench=False) + test_triton_batch_weighted_silu_and_mxfp8_quant(M=2048, N=2048, + n_experts=32, bench=False) + test_triton_batch_weighted_silu_and_mxfp8_quant(M=2048, N=1536, + n_experts=32, bench=False) + test_triton_batch_weighted_silu_and_mxfp8_quant(M=2048, N=1536, + n_experts=32, coef=10000.0, + grad_coef=10000.0, + bench=False) diff --git a/tests/test_smooth_quant.py b/tests/test_smooth_quant.py index b0011d9..1e0d45d 100644 --- a/tests/test_smooth_quant.py +++ b/tests/test_smooth_quant.py @@ -5,17 +5,18 @@ import torch +from linghe.facade.smooth_quant_linear import SmoothQuantLinear 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) + 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, +from linghe.tools.check import output_check +from linghe.tools.util import (torch_make_indices, torch_smooth_quant, round_up) -from linghe.facade.smooth_quant_linear import SmoothQuantLinear + def torch_split_smooth_quant(x_split, smooth_scales, round_scale=False): x_qs = [] @@ -97,17 +98,20 @@ def triton_split_smooth_quant(x_split, smooth_scales): def test_triton_smooth_quant(M=4096, N=4096, bench=False): device = 'cuda:0' x = torch.randn((M, N), dtype=torch.bfloat16, device=device) - smooth_scale = torch.randn((N,), device=device, dtype=torch.float32).abs() + 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, x_maxs_ref = torch_smooth_quant(x, smooth_scale, reverse=False, - round_scale=True) + round_scale=round_scale) 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') + reverse=False, + round_scale=round_scale, + calibrate=True) + output_check(x_q_ref, x_q, + 'triton_smooth_quant.data', rtol=rtol) output_check(scales_ref, x_scale, 'triton_smooth_quant.scale') output_check(x_maxs_ref, x_maxs, 'triton_smooth_quant.x_maxs') @@ -121,7 +125,7 @@ def test_triton_smooth_quant(M=4096, N=4096, bench=False): def test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, - size=16384): + size=16384): device = 'cuda:0' x = torch.randn((size,), dtype=torch.float32, device=device) x_q = torch.zeros((M, N), dtype=torch.bfloat16, device=device).to( @@ -141,8 +145,8 @@ def test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, 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) + 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') @@ -169,10 +173,10 @@ def test_triton_transpose_smooth_quant(M=4096, N=4096, bench=False): 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) + 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, @@ -196,7 +200,7 @@ def test_triton_transpose_smooth_quant(M=4096, N=4096, bench=False): def test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, - round_scale=False): + round_scale=False): device = 'cuda:0' P = round_up(M, b=32) y = torch.randn((M, N), dtype=torch.bfloat16, device=device) ** 3 @@ -211,8 +215,8 @@ def test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, 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) + reverse=True, + round_scale=round_scale) yt_gt, yt_scale_gt, yt_maxs_gt = torch_smooth_quant(y.t(), transpose_smooth_scale, @@ -225,12 +229,12 @@ def test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, 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) + 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 @@ -274,7 +278,8 @@ def test_triton_batch_smooth_quant(M=4096, N=4096, n_experts=32, topk=8, x_q_ref = torch.cat([x.view(torch.uint8) for x in x_q_ref], 0).view( torch.float8_e4m3fn) x_scale_ref = torch.cat(x_scale_ref, 0) - output_check(x_q_ref.float(), x_q.float(), 'triton_batch_smooth_quant.data') + 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') output_check(x_maxs_ref.float(), x_maxs.float(), @@ -294,30 +299,26 @@ def test_triton_batch_smooth_quant(M=4096, N=4096, n_experts=32, topk=8, n_repeat=n_repeat, ref_time=ref_time) - - def test_smooth_quant_linear(M=8192, N=1024, K=2048): - - dtype = torch.bfloat16 + 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') - 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') + output_check(dw_ref, dw, name='dw') if __name__ == '__main__': @@ -330,11 +331,11 @@ def test_smooth_quant_linear(M=8192, N=1024, K=2048): test_triton_smooth_quant(M=3457, N=512, bench=False) test_triton_subrow_smooth_quant(M=4096, N=5120, offset=5120, - size=2048) + size=2048) test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, - size=5120) + size=5120) test_triton_subrow_smooth_quant(M=4096, N=5120, offset=5120, - size=5120 * 10 - 1024) + 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) @@ -342,14 +343,14 @@ def test_smooth_quant_linear(M=8192, N=1024, K=2048): test_triton_transpose_smooth_quant(M=4096, N=3072, bench=False) test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, - round_scale=True) + round_scale=True) test_triton_transpose_rescale_smooth_quant(M=3895, N=4096, - round_scale=True) + round_scale=True) test_triton_transpose_rescale_smooth_quant(M=4096, N=3072, - round_scale=True) + round_scale=True) test_triton_transpose_rescale_smooth_quant(M=395, N=2048, - round_scale=True) + 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 + # test_smooth_quant_linear(M=8192, N=1024, K=2048) diff --git a/tests/test_topk.py b/tests/test_topk.py new file mode 100644 index 0000000..b5d5c7d --- /dev/null +++ b/tests/test_topk.py @@ -0,0 +1,211 @@ +# -*- 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..6cbaa96 100644 --- a/tests/test_transpose.py +++ b/tests/test_transpose.py @@ -8,16 +8,14 @@ import torch from linghe.tools.benchmark import benchmark_func -from linghe.tools.util import output_check +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) + 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() @@ -75,25 +73,36 @@ def test_nd_transpose(B=4096, M=4, N=4096, bench=False): 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) + 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] t_ref = torch_nd_transpose(x, 0, 1) - t = triton_transpose(x, dim0=0, dim1=1) + 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] t_ref = torch_nd_transpose(x, 0, 1) - t = triton_transpose(x, dim0=0, dim1=1) + 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, + 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) @@ -105,13 +114,7 @@ def test_transpose_and_pad(M=4095, N=4096, bench=False): dtype = torch.bfloat16 device = 'cuda:0' - n_repeat = 100 - - 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 @@ -127,7 +130,7 @@ def test_transpose_and_pad(M=4095, N=4096, bench=False): assert opt_output[:, -tail:].float().abs().sum().item() == 0 if bench: - benchmark_func(triton_transpose_and_pad, x_q, n_repeat=n_repeat, + benchmark_func(triton_transpose_and_pad, x_q, ref_bytes=M * N * 2) @@ -139,10 +142,12 @@ def test_batch_transpose(M=4096, N=4096, k=32, bench=False): torch.randn((M, N), dtype=dtype, device=device).to(torch.float8_e4m3fn) 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 @@ -159,12 +164,13 @@ def test_batch_transpose_and_pad(M=4096, N=4096, k=32, bench=False): xs = torch.randn((sum(count_list), N), dtype=dtype, device=device).to( 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 @@ -177,6 +183,6 @@ def test_batch_transpose_and_pad(M=4096, N=4096, k=32, bench=False): 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..b0938bb 100644 --- a/tests/test_unary.py +++ b/tests/test_unary.py @@ -3,29 +3,44 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import random + 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_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_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 +def test_calculate_smooth_scale(N=4096, bench=False): + 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 @@ -33,9 +48,47 @@ def test_calculate_smooth_scale(N=4096, bench=False): 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, + benchmark_func(torch_calculate_smooth_scale, x, n_repeat=n_repeat, ref_time=ref_time, ref_bytes=N * 8) + +def test_batch_clip(M=2048, N=1024, k=1024, clip_value=1.0, inf=False, + bench=False): + 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 bench: + 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_func(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_func(triton_batch_clip, xs4, clip_value, + ref_bytes=ref_bytes, ref_time=ref_time, + n_repeat=n_repeat, + n_warmup=0) + + if __name__ == '__main__': - test_calculate_smooth_scale(N=4096*32) - test_calculate_smooth_scale(N=4096*32-1897) + # test_calculate_smooth_scale(N=4096*32) + # test_calculate_smooth_scale(N=4096*32-1897) + # test_batch_clip(M=2048, N=8192, k=128, clip_value=0.1, bench=False) + # test_batch_clip(M=2048, N=1024, k=128, clip_value=1.0, bench=False) + test_batch_clip(M=2048, N=1024, k=128, clip_value=100.0, inf=True, + bench=False) From b67aad2566b14c243117fa6f504e92ab9bdf3939 Mon Sep 17 00:00:00 2001 From: "nanxiao.zy" Date: Wed, 14 Jan 2026 18:02:47 +0800 Subject: [PATCH 02/11] remove unused code --- scripts/test.sh | 15 ++- tests/test_dist_loss.py | 2 +- tests/test_gather.py | 105 +----------------- tests/test_la.py | 2 +- tests/test_loss.py | 22 ++-- tests/test_mla.py | 46 ++++---- tests/test_mxfp8_quant.py | 98 ----------------- tests/test_norm.py | 70 +----------- tests/test_reduce.py | 2 +- tests/test_rope.py | 2 +- tests/test_silu.py | 222 +------------------------------------- 11 files changed, 49 insertions(+), 537 deletions(-) delete mode 100644 tests/test_mxfp8_quant.py diff --git a/scripts/test.sh b/scripts/test.sh index 88316f7..bdaae0a 100644 --- a/scripts/test.sh +++ b/scripts/test.sh @@ -1,12 +1,12 @@ 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_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 && @@ -15,7 +15,6 @@ echo "test_hadamard_quant.py" && python test_hadamard_quant.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_mxfp8_quant.py" && python test_mxfp8_quant.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 && diff --git a/tests/test_dist_loss.py b/tests/test_dist_loss.py index 76d519b..b825c01 100644 --- a/tests/test_dist_loss.py +++ b/tests/test_dist_loss.py @@ -147,7 +147,7 @@ def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, grad_coef=1.0, 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) + inplace=False, group=pg, bench=False) 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) diff --git a/tests/test_gather.py b/tests/test_gather.py index c4f7e4e..035b3ae 100644 --- a/tests/test_gather.py +++ b/tests/test_gather.py @@ -6,14 +6,12 @@ import torch from linghe.quant.block import triton_batch_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 from linghe.tools.util import (torch_batch_smooth_quant, torch_blockwise_quant, torch_make_indices, - torch_smooth_quant, - torch_mxfp8_quant) + torch_smooth_quant) from linghe.utils.gather import (triton_make_row_id_map, triton_make_row_id_map_and_index, triton_index_select, @@ -23,7 +21,6 @@ triton_smooth_weighted_permute_with_indices, triton_batch_transpose_smooth_permute_with_indices, triton_batch_block_pad_permute_with_indices, - triton_batch_mxfp8_permute_with_indices ) @@ -196,51 +193,6 @@ def torch_batch_block_pad_permute_with_indices(x, 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] - - 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) - 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 test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): dtype = torch.bfloat16 device = 'cuda:0' @@ -576,54 +528,6 @@ def test_batch_block_pad_permute_with_indices(M=16384, N=2048, n_experts=32, ref_bytes=num_out_tokens * N * 4) -def test_batch_mxfp8_permute_with_indices(M=16384, N=2048, n_experts=32, topk=2, - bench=False): - 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_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, - indices, - token_count_per_expert_list, - probs=probs) - 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') - - if bench: - num_out_tokens = sum(token_count_per_expert_list) - benchmark_func(triton_batch_mxfp8_permute_with_indices, x, - token_count_per_expert, - indices, - token_count_per_expert_list, - probs=probs, - ref_bytes=num_out_tokens * N * 4) - - benchmark_func(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_func(triton_batch_mxfp8_quant, xs, token_count_per_expert, - token_count_per_expert_list, - ref_bytes=num_out_tokens * N * 4) - if __name__ == '__main__': test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False) @@ -653,10 +557,3 @@ def test_batch_mxfp8_permute_with_indices(M=16384, N=2048, n_experts=32, topk=2, bench=False) test_batch_block_pad_permute_with_indices(M=8192, N=1536, n_experts=32, topk=2, bench=False) - - test_batch_mxfp8_permute_with_indices(M=3095, N=2048, n_experts=32, topk=2, - bench=False) - test_batch_mxfp8_permute_with_indices(M=0, N=2048, n_experts=32, topk=2, - bench=False) - test_batch_mxfp8_permute_with_indices(M=1024, N=1536, n_experts=32, topk=2, - bench=False) diff --git a/tests/test_la.py b/tests/test_la.py index 76e7a9f..6661e32 100644 --- a/tests/test_la.py +++ b/tests/test_la.py @@ -197,4 +197,4 @@ def test_la(bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, if __name__ == '__main__': test_la(bs=1, length=8192, qo_heads=64, kv_heads=64, dim=128, digest=False, - bench=True) + bench=False) diff --git a/tests/test_loss.py b/tests/test_loss.py index 256f61b..a95fb1b 100644 --- a/tests/test_loss.py +++ b/tests/test_loss.py @@ -137,32 +137,32 @@ def test_z_loss(L=4096, B=2, N=256, coef=0.001, bench=False): if __name__ == '__main__': test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, - inplace=True, bench=True) + inplace=True, bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, - grad_coef=1e-6, inplace=True, bench=True) + grad_coef=1e-6, inplace=True, bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=10000.0, grad_coef=100.0, fill=True, inplace=True, - bench=True) + bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, fill=True, ignore_index=-100, - inplace=True, bench=True) + inplace=True, bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, fill=True, ignore_index=0, inplace=True, - bench=True) + bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184 - 16, coef=10000.0, grad_coef=100.0, fill=True, inplace=True, - bench=True) + bench=False) test_triton_softmax_cross_entropy(M=8192, N=175175, coef=1.0, grad_coef=1.0, - inplace=True, bench=True) + inplace=True, bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=0.0, grad_coef=0.0, - inplace=True, bench=True) + inplace=True, bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=0.0, - grad_coef=100.0, inplace=True, bench=True) + grad_coef=100.0, inplace=True, bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1000.0, - grad_coef=0.0, inplace=True, bench=True) + grad_coef=0.0, inplace=True, bench=False) test_triton_softmax_cross_entropy(M=8192, N=157184, coef=100.0, grad_coef=100.0, fill=True, inplace=True, - bench=True) + bench=False) test_triton_softmax_cross_entropy(M=4096, N=157184, coef=0.1, grad_coef=1.0, inplace=True, bench=False) diff --git a/tests/test_mla.py b/tests/test_mla.py index a213863..74c8a74 100644 --- a/tests/test_mla.py +++ b/tests/test_mla.py @@ -324,29 +324,29 @@ def test_fp8_mla(B=2, L=4096, H=16, causal=True, hpc=False, quant_value=False, if __name__ == "__main__": - # test_softmax(M=128, N=128) + test_softmax(M=128, N=128) - # test_dot_sum(M=128, N=128, D=128) + test_dot_sum(M=128, N=128, D=128) test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=False, coef=1.0, - clip_value=500, bench=True) - # test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=500.0, bench=True) - # test_mla(B=1, L=8192, H=64, causal=True, hpc=True, safe=False, coef=1.0, clip_value=None, bench=False) - # test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=True, coef=100.0, clip_value=None, bench=False) - # test_mla(B=1, L=4096, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - # test_mla(B=1, L=4096, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - # test_mla(B=1, L=8192, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - # test_mla(B=1, L=8192, H=1, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - - # test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, pad=False, bench=True) - # test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=100.0, pad=False, bench=True) - # test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=None, pad=True, bench=True) - - # test_varlen_mla(LS=[4096,4096], H=64, causal=True, hpc=False, safe=True, coef=1.0, bench=False) - # test_varlen_mla(LS=[2048,2048,4096], H=64, causal=True, hpc=True, safe=True, coef=1.0, bench=False) - # test_varlen_mla(LS=[127,873,3096], H=64, causal=False, hpc=False, safe=False, coef=1.0, bench=False) - # test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, clip_value=100.0, bench=False) - # test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, pad=True, bench=False) - # test_varlen_mla(LS=[127,873,3456], H=1, causal=True, hpc=False, safe=True, coef=1.0, bench=False) - - # test_fp8_mla(B=1, L=8192, H=64, causal=True, hpc=False, quant_value=False, bench=True) + clip_value=500, bench=False) + test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=500.0, bench=False) + test_mla(B=1, L=8192, H=64, causal=True, hpc=True, safe=False, coef=1.0, clip_value=None, bench=False) + test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=True, coef=100.0, clip_value=None, bench=False) + test_mla(B=1, L=4096, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + test_mla(B=1, L=4096, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + test_mla(B=1, L=8192, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + test_mla(B=1, L=8192, H=1, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) + + test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, pad=False, bench=False) + test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=100.0, pad=False, bench=False) + test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=None, pad=True, bench=False) + + test_varlen_mla(LS=[4096,4096], H=64, causal=True, hpc=False, safe=True, coef=1.0, bench=False) + test_varlen_mla(LS=[2048,2048,4096], H=64, causal=True, hpc=True, safe=True, coef=1.0, bench=False) + test_varlen_mla(LS=[127,873,3096], H=64, causal=False, hpc=False, safe=False, coef=1.0, bench=False) + test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, clip_value=100.0, bench=False) + test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, pad=True, bench=False) + test_varlen_mla(LS=[127,873,3456], H=1, causal=True, hpc=False, safe=True, coef=1.0, bench=False) + + test_fp8_mla(B=1, L=8192, H=64, causal=True, hpc=False, quant_value=False, bench=False) diff --git a/tests/test_mxfp8_quant.py b/tests/test_mxfp8_quant.py deleted file mode 100644 index 2054306..0000000 --- a/tests/test_mxfp8_quant.py +++ /dev/null @@ -1,98 +0,0 @@ -# -*- coding: utf-8 -*- -""" -Copyright (c) Ant Financial Service Group and its affiliates. -""" - -import random - -import torch - -from linghe.quant.mxfp8 import triton_mxfp8_quant, triton_batch_mxfp8_quant -from linghe.tools.benchmark import benchmark_func -from linghe.tools.check import output_check -from linghe.tools.util import torch_mxfp8_quant - - -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 test_mxfp8_quant(M=4096, N=4096, bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' - - x = torch.randn(M, N, dtype=dtype, device=device) - - x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = torch_mxfp8_quant(x) - x_q, x_scale, xt_q, xt_scale = triton_mxfp8_quant(x) - - 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 * 4 - benchmark_func(triton_mxfp8_quant, x, ref_bytes=ref_bytes) - benchmark_func(torch_mxfp8_quant, x) - - -def test_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] - token_count_per_expert = torch.tensor(splits, device=device) - - x = torch.randn((sum(splits), N), dtype=dtype, device=device) - - x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = torch_batch_mxfp8_quant(x, - splits) - - 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(torch_batch_mxfp8_quant, x, splits) - benchmark_func(triton_batch_mxfp8_quant, x, token_count_per_expert, - splits, output_mode=2, ref_bytes=ref_bytes) - - -if __name__ == '__main__': - test_mxfp8_quant(M=4096, N=8192, bench=False) - test_mxfp8_quant(M=4031, N=8192, bench=False) - test_mxfp8_quant(M=4096, N=8192, bench=False) - test_batch_mxfp8_quant(M=4096, N=8192, bench=False) diff --git a/tests/test_norm.py b/tests/test_norm.py index 61bb916..e3359aa 100644 --- a/tests/test_norm.py +++ b/tests/test_norm.py @@ -10,11 +10,9 @@ 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, - torch_mxfp8_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_and_mxfp8_quant_forward, triton_rms_norm_fp32_gemm_block_quant_forward, triton_rms_norm_backward, triton_rms_norm_forward) @@ -129,24 +127,6 @@ def split_rms_gemm_block_quant_forward(x, norm_weight, route_weight, return y, logit, q, s -def torch_rms_and_mxfp8_quant_forward(x, weight): - 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 - ) - with torch.no_grad(): - rmsnorm.weight.copy_(weight) - y = rmsnorm(x) - # mxfp8 - y_q, y_scale, yt_q, yt_scale = torch_mxfp8_quant(y) - return y_q, y_scale, yt_q, yt_scale - - def test_rmsnorm(M=4096, N=4096, bench=False): dtype = torch.bfloat16 device = 'cuda:0' @@ -267,49 +247,6 @@ def test_rmsnorm_and_block_quant(M=4096, N=4096, bench=False): ref_bytes=M * N * 6) -def test_rmsnorm_and_mxfp8_quant(M=4096, N=4096, bench=False): - dtype = torch.bfloat16 - 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) - - # mxfp8 - q_ref, scale_ref, qt_ref, scale_t_ref = torch_rms_and_mxfp8_quant_forward(x, - weight) - q, scale, rms, q_t, scale_t = triton_rms_norm_and_mxfp8_quant_forward(x, - weight, - 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") - - q, scale, _, _, _ = triton_rms_norm_and_mxfp8_quant_forward(x, weight, - 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_mxfp8_quant_forward(x, weight, - rms=rms, - output_mode=1) - output_check(qt_ref, q_t, name='0.block.t_data', rtol=0.125) - output_check(scale_t_ref, scale_t, name="0.block.t_scale") - - if bench: - benchmark_func(triton_rms_norm_and_mxfp8_quant_forward, x, weight, - output_mode=0, - ref_bytes=M * N * 3) - - benchmark_func(triton_rms_norm_and_mxfp8_quant_forward, x, weight, - output_mode=1, - ref_bytes=M * N * 3) - - benchmark_func(triton_rms_norm_and_mxfp8_quant_forward, x, weight, - output_mode=2, - ref_bytes=M * N * 4) - - def test_rms_norm_fp32_gemm_block_quant_forward(M=8192, N=256, K=2048, bench=False): dtype = torch.bfloat16 @@ -378,10 +315,7 @@ def test_rms_norm_fp32_gemm_block_quant_forward(M=8192, N=256, K=2048, 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=True) - - test_rmsnorm_and_mxfp8_quant(M=2048, N=1664, bench=False) - test_rmsnorm_and_mxfp8_quant(M=8192, N=4096, 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) diff --git a/tests/test_reduce.py b/tests/test_reduce.py index 59617e8..82506f1 100644 --- a/tests/test_reduce.py +++ b/tests/test_reduce.py @@ -119,4 +119,4 @@ def test_batch_norm(M=4096, N=8192, k=32, coef=1.0, bench=False): test_norm(M=100000, N=8192, bench=False) test_batch_norm(M=4096, N=1024, k=16, bench=False) test_batch_norm(M=4096, N=1024, k=64, bench=False) - test_batch_norm(M=4096, N=2048, k=1024, coef=1e12, bench=True) + test_batch_norm(M=4096, N=2048, k=1024, coef=1e12, bench=False) diff --git a/tests/test_rope.py b/tests/test_rope.py index 61b0b55..4ed8994 100644 --- a/tests/test_rope.py +++ b/tests/test_rope.py @@ -653,7 +653,7 @@ def test_varlen_mla_rope(lengths=[2048, 2048], H=32, rope_theta=10000.0, transposed=True, silu=False, 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=False, silu=True, bench=True) + transposed=False, silu=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=False, silu=False, bench=False) diff --git a/tests/test_silu.py b/tests/test_silu.py index fac5c88..3713f45 100644 --- a/tests/test_silu.py +++ b/tests/test_silu.py @@ -16,18 +16,13 @@ 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) + torch_group_quant) from linghe.tools.check import output_check @@ -88,16 +83,6 @@ def torch_silu_and_block_quant_forward(x, round_scale=True): return y_q, y_scale, yt_q, yt_scale -def torch_silu_and_mxfp8_quant_forward(x): - 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, 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, round_scale=True, reverse=True): @@ -129,16 +114,6 @@ def torch_silu_and_block_quant_backward(grad, x, round_scale=True): return q, dx_scale, yt_q, yt_scale -def torch_silu_and_mxfp8_quant_backward(grad, x): - grad = grad.float() - x = x.float().detach().clone().requires_grad_() - y = torch_silu(x) - 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, @@ -217,44 +192,6 @@ def torch_batch_weighted_silu_and_block_quant_forward(xs, weight, return qs, scales, qts, qtscales -def torch_batch_weighted_silu_and_mxfp8_quant_forward(xs, weight, - counts): - 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]) - - 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, @@ -351,44 +288,6 @@ def torch_batch_weighted_silu_and_block_quant_backward(grad_output, x, weight, return dx_q, dx_scale, dw, qts, qtscales -def torch_batch_weighted_silu_and_mxfp8_quant_backward(grad_output, x, weight, - counts): - 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) - 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 - def test_weighted_silu(M=4096, N=4096, asm=False, coef=1.0, bench=False): x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') @@ -525,48 +424,6 @@ def test_silu_and_block_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, n_repeat=100, ref_bytes=M * N * 5) -def test_silu_and_mxfp8_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, - bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') - x = (x * coef).clone().detach().requires_grad_() - grad_output = torch.randn((M, N // 2), 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) - y_q, y_scale, yt_q, yt_scale = triton_silu_and_mxfp8_quant_forward(x, - 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, - 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, - 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) - dx_q, dx_scale, dxt_q, dxt_scale = triton_silu_and_mxfp8_quant_backward( - grad_output, x) - 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') - - if bench: - benchmark_func(triton_silu_and_mxfp8_quant_forward, x, - n_repeat=100, ref_bytes=M * N * 3) - benchmark_func(triton_silu_and_mxfp8_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, coef=1.0, @@ -735,68 +592,6 @@ def test_triton_batch_weighted_silu_and_block_quant(M=1024, N=4096, ref_bytes=n_experts * M * N * 4, ref_time=ref_time) -def test_triton_batch_weighted_silu_and_mxfp8_quant(M=1024, N=4096, - n_experts=32, - coef=1.0, - grad_coef=1.0, - bench=False): - 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') * coef - weight = torch.randn((bs, 1), dtype=torch.float32, device='cuda:0') - - grad_output = torch.randn((bs, N // 2), 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) - x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_mxfp8_quant_forward( - x, - weight, - counts, - count_list, - 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) - dx, dx_scale, dw, dxt, dxt_scale = triton_batch_weighted_silu_and_mxfp8_quant_backward( - grad_output, x, weight, counts, splits=count_list) - 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) - - if bench: - ref_time = None - benchmark_func(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 * 2.5, ref_time=ref_time) - benchmark_func(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 * 3, ref_time=ref_time) - benchmark_func(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 * 4, ref_time=ref_time) - if __name__ == '__main__': test_weighted_silu(M=16384, N=4096, coef=1.0, asm=False, bench=False) @@ -818,10 +613,6 @@ def test_triton_batch_weighted_silu_and_mxfp8_quant(M=1024, N=4096, test_silu_and_block_quant(M=4096, N=1536 * 8, coef=0.0, grad_coef=0.0, bench=False) - test_silu_and_mxfp8_quant(M=16384, N=1024, bench=False) - test_silu_and_mxfp8_quant(M=2345, N=1024, bench=False) - test_silu_and_mxfp8_quant(M=2345, N=1536, 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_smooth_quant(M=2048, N=2048, @@ -838,14 +629,3 @@ def test_triton_batch_weighted_silu_and_mxfp8_quant(M=1024, N=4096, test_triton_batch_weighted_silu_and_block_quant(M=12080, N=1536, n_experts=32, coef=0.0, grad_coef=0.0, bench=False) - - test_triton_batch_weighted_silu_and_mxfp8_quant(M=0, N=2048, n_experts=32, - bench=False) - test_triton_batch_weighted_silu_and_mxfp8_quant(M=2048, N=2048, - n_experts=32, bench=False) - test_triton_batch_weighted_silu_and_mxfp8_quant(M=2048, N=1536, - n_experts=32, bench=False) - test_triton_batch_weighted_silu_and_mxfp8_quant(M=2048, N=1536, - n_experts=32, coef=10000.0, - grad_coef=10000.0, - bench=False) From e8ad21b30f517e9ce992a31f59d91a1fe1cef1c7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8D=97=E9=9C=84?= Date: Wed, 14 Jan 2026 18:17:47 +0800 Subject: [PATCH 03/11] 0.3.0 --- docs/linghe.html | 3 + docs/linghe/attn.html | 238 ++++++++ docs/linghe/attn/la.html | 239 ++++++++ docs/linghe/attn/mla.html | 239 ++++++++ docs/linghe/experimental.html | 246 ++++++++ docs/linghe/experimental/demb.html | 303 ++++++++++ docs/linghe/experimental/dla.html | 239 ++++++++ docs/linghe/experimental/dmm.html | 271 +++++++++ .../gmem_barrier_arrive_wait.html | 272 +++++++++ .../linghe/experimental/symm_mem_barrier.html | 288 ++++++++++ docs/linghe/experimental/test_demb.html | 239 ++++++++ docs/linghe/experimental/test_dla.html | 239 ++++++++ docs/linghe/experimental/test_dmm.html | 239 ++++++++ docs/linghe/facade.html | 6 + docs/linghe/facade/emb.html | 336 +++++++++++ docs/linghe/facade/gate.html | 273 +++++++++ docs/linghe/facade/loss.html | 33 +- docs/linghe/facade/mla.html | 270 +++++++++ docs/linghe/facade/norm.html | 151 ++++- docs/linghe/facade/permutation.html | 299 ++++++++++ docs/linghe/facade/rope.html | 61 +- docs/linghe/facade/silu.html | 543 ++++++++++++++++++ docs/linghe/facade/topk.html | 310 ++++++++++ docs/linghe/facade/transpose.html | 12 +- docs/linghe/gemm/fp32_gemm.html | 88 +-- docs/linghe/quant/block.html | 78 ++- docs/linghe/tools.html | 239 ++++++++ docs/linghe/tools/benchmark.html | 239 ++++++++ docs/linghe/tools/check.html | 239 ++++++++ docs/linghe/tools/util.html | 239 ++++++++ docs/linghe/utils.html | 6 +- docs/linghe/utils/emb.html | 366 ++++++++++++ docs/linghe/utils/gate.html | 272 +++++++++ docs/linghe/utils/gather.html | 51 +- docs/linghe/utils/loss.html | 139 ++++- docs/linghe/utils/mul.html | 334 +++++++++++ docs/linghe/utils/norm.html | 44 +- docs/linghe/utils/rearange.html | 10 +- docs/linghe/utils/reduce.html | 51 +- docs/linghe/utils/rope.html | 163 +++++- docs/linghe/utils/silu.html | 2 +- docs/linghe/utils/topk.html | 368 ++++++++++++ docs/linghe/utils/transpose.html | 5 +- docs/linghe/utils/unary.html | 271 +++++++++ docs/search.js | 2 +- linghe/experimental/norm.py | 352 ------------ linghe/experimental/test_norm.py | 87 --- scripts/dev.py | 30 - 48 files changed, 8440 insertions(+), 584 deletions(-) create mode 100644 docs/linghe/attn.html create mode 100644 docs/linghe/attn/la.html create mode 100644 docs/linghe/attn/mla.html create mode 100644 docs/linghe/experimental.html create mode 100644 docs/linghe/experimental/demb.html create mode 100644 docs/linghe/experimental/dla.html create mode 100644 docs/linghe/experimental/dmm.html create mode 100644 docs/linghe/experimental/gmem_barrier_arrive_wait.html create mode 100644 docs/linghe/experimental/symm_mem_barrier.html create mode 100644 docs/linghe/experimental/test_demb.html create mode 100644 docs/linghe/experimental/test_dla.html create mode 100644 docs/linghe/experimental/test_dmm.html create mode 100644 docs/linghe/facade/emb.html create mode 100644 docs/linghe/facade/gate.html create mode 100644 docs/linghe/facade/mla.html create mode 100644 docs/linghe/facade/permutation.html create mode 100644 docs/linghe/facade/silu.html create mode 100644 docs/linghe/facade/topk.html create mode 100644 docs/linghe/tools.html create mode 100644 docs/linghe/tools/benchmark.html create mode 100644 docs/linghe/tools/check.html create mode 100644 docs/linghe/tools/util.html create mode 100644 docs/linghe/utils/emb.html create mode 100644 docs/linghe/utils/gate.html create mode 100644 docs/linghe/utils/mul.html create mode 100644 docs/linghe/utils/topk.html create mode 100644 docs/linghe/utils/unary.html delete mode 100644 linghe/experimental/norm.py delete mode 100644 linghe/experimental/test_norm.py delete mode 100644 scripts/dev.py 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:
    + +
      +
    • logits: logits tensor, shape [...,dim]
    • +
    • coef: z loss coef
    • +
    + +
    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:
      -
    • 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]
    • +
    +
    +
    + + +
    +
    +
    + + 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:
    + +
      +
    • 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]
    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/experimental/norm.py b/linghe/experimental/norm.py deleted file mode 100644 index d69b66b..0000000 --- a/linghe/experimental/norm.py +++ /dev/null @@ -1,352 +0,0 @@ -# -*- 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 - -""" -the code is used to reproduce the barrier bug in cross-block reduce. -""" - - -@triton.jit -def rms_norm_forward_kernel(x_ptr, - weight_ptr, - out_ptr, - cache_ptr, - signal_ptr, - rms_ptr, - eps, - M, - n: tl.constexpr, - H: tl.constexpr, - B: tl.constexpr): - rid = tl.program_id(axis=0) - cid = tl.program_id(axis=1) - CB = tl.num_programs(1) - - indices = rid * H + tl.arange(0, H) - offs = rid * H * n + cid * B + tl.arange(0, H)[:, None] * n + tl.arange(0, - B)[ - None, :] - - x = tl.load(x_ptr + offs).to(tl.float32) - weight = tl.load(weight_ptr + cid * B + tl.arange(0, B)).to(tl.float32) - - s = tl.sum(x * x, axis=1) - tl.atomic_add(cache_ptr + indices, s, sem='acq_rel', scope='sys') - # tl.debug_barrier() - tl.atomic_add(signal_ptr + rid, 1, sem='acq_rel', scope='sys') - tl.debug_barrier() - # tl.inline_asm_elementwise( - # "membar.gl;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 - # ) - # tl.inline_asm_elementwise( - # "bar.sync 0;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 - # ) - count = tl.load(signal_ptr + rid, cache_modifier='.cv') - while count < CB: - count = tl.load(signal_ptr + rid, cache_modifier='.cv') - # if cid + rid == 0: - # tl.device_print('count', count) - tl.debug_barrier() - - sums = tl.load(cache_ptr + indices, cache_modifier='.cv', volatile=True) - # sums = tl.atomic_add(cache_ptr + indices, 0.0, sem='acq_rel', scope='gpu') - - rms = tl.rsqrt(sums / n + eps) - if cid == 0: - tl.store(rms_ptr + indices, rms) - - x = x * rms[:, None] * weight[None, :] - - tl.store(out_ptr + offs, x) - - -def triton_rms_norm_forward(x: torch.Tensor, - weight: torch.Tensor, - eps: float = 1e-6): - """ - Fused RMSNorm forward. - Args: - x: Input tensor, shape [M, N] - weight: RMSNorm weight, shape [N] - eps: epsilon value for L2 normalization. - Returns: - - out: output. - - rms: Reciprocal of the root mean square of the - input calculated over the last dimension. - """ - assert x.is_contiguous() and weight.is_contiguous() - M, n = x.shape - device = x.device - - out = torch.empty((M, n), dtype=x.dtype, device=device) - rms = torch.empty((M,), dtype=torch.float32, device=device) - H = 8 - B = 128 - assert M % H == 0 and n % B == 0 - CB = n // 128 # column block - RB = M // H # row block - cache = torch.zeros((M,), dtype=torch.float32, device=device) - signals = torch.zeros((RB,), dtype=torch.int32, device=device) - grid = (RB, CB) - rms_norm_forward_kernel[grid]( - x, - weight, - out, - cache, - signals, - rms, - eps, - M, - n, - H, - B, - num_stages=1, - num_warps=1 - ) - return out, rms - - -@triton.jit -def _parallel_rms_norm_and_block_quant_forward_kernel(x_ptr, - weight_ptr, - out_ptr, - scale_ptr, - transpose_output_ptr, - transpose_scale_ptr, - cache_ptr, - rms_ptr, - eps, - M, - n: tl.constexpr, - H: tl.constexpr, - B: tl.constexpr, - K: tl.constexpr, - ROUND: tl.constexpr): - rid = tl.program_id(axis=0) - cid = tl.program_id(axis=1) - CB = tl.num_programs(1) - - indices = rid * H + tl.arange(0, H) - masks = indices[:, None] < M - offs = rid * H * n + cid * B + tl.arange(0, H)[:, None] * n + tl.arange(0, - B)[ - None, :] - - x = tl.load(x_ptr + offs, mask=masks).to(tl.float32) - s = tl.sum(x * x, axis=1) - # tl.debug_barrier() - tl.atomic_add(cache_ptr + indices, s, scope='sys') - # tl.debug_barrier() - tl.atomic_add(cache_ptr + M + rid, 1.0, scope='sys') - # tl.inline_asm_elementwise( - # "bar.sync 0;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 - # ) - tl.debug_barrier() - # count = tl.atomic_add(cache_ptr + M + rid, 0.0) - # tl.inline_asm_elementwise( - # "membar.gl;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 - # ) - for i in range(3): - count = tl.load(cache_ptr + M + rid) - # tl.debug_barrier() - while count < CB: - count = tl.load(cache_ptr + M + rid) - tl.debug_barrier() - # tl.inline_asm_elementwise( - # "membar.gl;", "=r", [], dtype=tl.int32, is_pure=False, pack=1 - # ) - sums = tl.load(cache_ptr + indices, mask=indices < M) - - rms = tl.rsqrt(sums / n + eps) - if cid == CB - 1: - tl.store(rms_ptr + indices, rms, mask=indices < M) - - toffs = cid * M * B + rid * H + tl.arange(0, B)[:, None] * M + tl.arange(0, - H)[ - None, :] - weight = tl.load(weight_ptr + cid * B + tl.arange(0, B)).to(tl.float32) - x = x * rms[:, None] * weight[None, :] - scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - q = (x / scale[:, None]).to(out_ptr.dtype.element_ty) - - tl.store(scale_ptr + cid * M + indices, scale, mask=indices < M) - tl.store(out_ptr + offs, q, mask=masks) - - 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 * B + tl.arange(0, B), scale) - - q = (tl.trans(x / scale)).to(transpose_output_ptr.dtype.element_ty) - tl.store(transpose_output_ptr + toffs, q, mask=indices[None, :] < M) - - -@triton.jit -def parallel_rms_norm_and_block_quant_forward_kernel(x_ptr, - weight_ptr, - out_ptr, - scale_ptr, - transpose_output_ptr, - transpose_scale_ptr, - cache_ptr, - rms_ptr, - eps, - M, - n: tl.constexpr, - H: tl.constexpr, - B: tl.constexpr, - K: tl.constexpr, - ROUND: tl.constexpr): - rid = tl.program_id(axis=0) - cid = tl.program_id(axis=1) - CB = tl.num_programs(1) - - indices = rid * H + tl.arange(0, H) - masks = indices[:, None] < M - offs = rid * H * n + cid * K * B + tl.arange(0, H)[:, None] * n + tl.arange( - 0, B)[ - None, :] - - s = tl.zeros((H,), dtype=tl.float32) - for i in range(K): - x = tl.load(x_ptr + i * B + offs, mask=masks).to(tl.float32) - s += tl.sum(x * x, axis=1) - - tl.atomic_add(cache_ptr + indices, s) - tl.atomic_add(cache_ptr + M + rid, 1.0) - tl.debug_barrier() - - count = tl.atomic_add(cache_ptr + M + rid, 0.0) - # count = tl.load(cache_ptr + M + rid) - tl.debug_barrier() - while count < 1.0 * CB: - count = tl.load(cache_ptr + M + rid) - for i in range(1024): - pass - tl.debug_barrier() - sums = tl.load(cache_ptr + indices, mask=indices < M) - - rms = tl.rsqrt(sums / n + eps) - - if cid == CB - 1: - tl.store(rms_ptr + indices, rms, mask=indices < M) - - toffs = cid * M * B * K + rid * H + tl.arange(0, B)[:, - None] * M + tl.arange(0, H)[ - None, :] - for i in range(K): - weight = tl.load(weight_ptr + cid * K * B + i * B + tl.arange(0, B)).to( - tl.float32) - x = tl.load(x_ptr + i * B + offs, mask=masks).to(tl.float32) - x = x * rms[:, None] * weight[None, :] - scale = tl.maximum(tl.max(tl.abs(x), 1) / 448.0, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) - q = (x / scale[:, None]).to(out_ptr.dtype.element_ty) - - tl.store(scale_ptr + cid * K * M + i * M + indices, scale, - mask=indices < M) - tl.store(out_ptr + i * B + offs, q, mask=masks) - - 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 * B * K + i * B + tl.arange(0, - B), - scale) - - q = (tl.trans(x / scale)).to(transpose_output_ptr.dtype.element_ty) - tl.store(transpose_output_ptr + i * B * M + toffs, q, - mask=indices[None, :] < M) - - -def triton_parallel_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: - x: Input tensor, shape [M, N] - weight: RMSNorm weight, shape [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. - output_mode: one of {0, 1, 2}. - 0: only output non-transpose tensor - 1: only output transposed tensor - 2: return both - Returns: - - out: quantization data. - - scale: quantization scale. - - rms: Reciprocal of the root mean square of the - input calculated over the last dimension. - - transpose_output: quantization data of transposed gradient. - - transpose_scale: quantization scale of transposed gradient. - """ - # row-wise read, row-wise write - assert x.is_contiguous() and weight.is_contiguous() - M, n = x.shape - 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((n // 128, M), device=device, dtype=torch.float32) - - 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) - - assert rms is None - rms = torch.empty((M,), dtype=torch.float32, device=device) - H = 128 - B = 128 - CB = n // B - K = n // (B * CB) - assert K >= 1 - RB = triton.cdiv(M, H) - cache = torch.zeros((M + RB,), dtype=torch.float32, device=device) - grid = (RB, CB) - _parallel_rms_norm_and_block_quant_forward_kernel[grid]( - x, - weight, - out, - scale, - transpose_output, - transpose_scale, - cache, - rms, - eps, - M, - n, - H, - B, - K, - round_scale, - num_stages=5, - num_warps=4 - ) - return out, scale, rms, transpose_output, transpose_scale diff --git a/linghe/experimental/test_norm.py b/linghe/experimental/test_norm.py deleted file mode 100644 index 95074f8..0000000 --- a/linghe/experimental/test_norm.py +++ /dev/null @@ -1,87 +0,0 @@ -# -*- coding: utf-8 -*- -""" -Copyright (c) Ant Financial Service Group and its affiliates. -""" - -import torch - -from linghe.experimental.norm import triton_rms_norm_forward, \ - triton_parallel_rms_norm_and_block_quant_forward -from linghe.tools.benchmark import benchmark_func -from linghe.tools.check import output_check -from linghe.utils.norm import triton_rms_norm_and_block_quant_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 - ) - with torch.no_grad(): - rmsnorm.weight.copy_(weight) - rms = torch.rsqrt(torch.sum(x ** 2, 1) / N + 1e-6) - return rmsnorm(x).to(dtype), rms - - -def test_norm(M=4096, N=4096, bench=False): - dtype = torch.bfloat16 - device = 'cuda:0' - - x = torch.ones(M, N, dtype=dtype, requires_grad=False, device=device) - weight = torch.ones(N, dtype=dtype, requires_grad=False, device=device) - - output_ref, rms_ref = torch_rms_forward(x, weight) - output, rms = triton_rms_norm_forward(x, weight) - - output_check(rms_ref, rms, name="rms", rtol=0.001) - output_check(output_ref, output, name="output", rtol=0.001) - - -def test_parallel_rmsnorm_and_block_quant(M=4096, N=4096, bench=False): - dtype = torch.bfloat16 - 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) - - q_ref, scale_ref, rms_ref, qt_ref, scale_t_ref = triton_rms_norm_and_block_quant_forward( - x, weight, - round_scale=False, - output_mode=2) - - q, scale, rms, qt, scale_t = triton_parallel_rms_norm_and_block_quant_forward( - x, weight, - round_scale=False, - output_mode=2) - output_check(q_ref, q, name='parallel.block.data', rtol=-0.125) - output_check(scale_ref, scale, name="parallel.block.scale", rtol=-0.125) - output_check(rms_ref, rms, name="parallel.block.rms", rtol=-0.125) - output_check(qt_ref, qt, name='parallel.block.t_data', rtol=-0.125) - output_check(scale_t_ref, scale_t, name="parallel.block.t_scale", - rtol=-0.125) - - if bench: - benchmark_func(triton_rms_norm_and_block_quant_forward, x, weight, - round_scale=False, - output_mode=2, - ref_bytes=M * N * 4, - n_profile=2) - - benchmark_func(triton_parallel_rms_norm_and_block_quant_forward, x, - weight, - round_scale=False, - output_mode=2, - ref_bytes=M * N * 4, - n_profile=2) - - -if __name__ == '__main__': - # /usr/local/lib/python3.12/dist-packages/triton/backends/nvidia/bin/ptxas -lineinfo -v --gpu-name=sm_90a /tmp/tmp3l_m5rfp.ptx -o /tmp/tmp3l_m5rfp.ptx.o - test_norm(M=1024, N=4096, bench=False) - # test_parallel_rmsnorm_and_block_quant(M=4096, N=4096, bench=False) diff --git a/scripts/dev.py b/scripts/dev.py deleted file mode 100644 index b46d113..0000000 --- a/scripts/dev.py +++ /dev/null @@ -1,30 +0,0 @@ -import torch -import triton -import triton.language as tl - - -def test_cpu_gpu_diff(): - x = torch.randn((128, 128), dtype=torch.float32) * 100 - cos = torch.cos(x) - c = torch.cos(x.cuda()) - torch.testing.assert_close(cos, c.cpu()) - - -@triton.jit -def index_overflow(x): - i = tl.program_id(0) - ptr = x + i * 2 ** 30 - offs = 2 ** 30 + 2 ** 30 + 2 ** 30 + 2 ** 30 - # offs = (i).to(tl.int64) * 2 ** 30 + 5*2**32 - tl.store(x + i, offs) - - -def test_index_overflow(): - x = torch.zeros((128,), dtype=torch.int64, device='cuda:0') - index_overflow[(128,)](x) - print(x) - - -if __name__ == '__main__': - # test_cpu_gpu_diff() - test_index_overflow() From b596bd4b35da9fa43ca34a5061b4a7a6653eaeee Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8D=97=E9=9C=84?= Date: Wed, 14 Jan 2026 18:25:53 +0800 Subject: [PATCH 04/11] remove unused files --- .aci.yml | 86 -------------------------------------------------------- 1 file changed, 86 deletions(-) delete mode 100644 .aci.yml diff --git a/.aci.yml b/.aci.yml deleted file mode 100644 index f3dceb3..0000000 --- a/.aci.yml +++ /dev/null @@ -1,86 +0,0 @@ -version: "2.0" - -stages: -- 前置检查 -- 构建&发布到测试库 -- 验包&确认 -- 发布正式库前检查 -- 发布正式库 - -jobs: - 单元测试: - stage: 前置检查 - component: python-ut - inputs: - languageConfig: - pythonVersion: "3.9.0" - config: - execute: - isAllowSkip: true - - 代码检查: - stage: 前置检查 - component: python-sast - inputs: - excludes: # 选填项,排除哪些项不进行代码扫描 - - "**__init__.py**" - - "**/tests/**" - - ansible/* - - config/* - - benchmark/* - - scripts/* - - examples/* - - docs/* - codePath: "./" # 选填项,选择扫描目录 - config: - execute: - timeout: 600 # 选填项,任务超时时间 - isAllowSkip: true - afterExecute: - checkRule: # 选填项,卡点策略, 根据实际团队质量要求进行配置 - - ${{outputs.critical}} <= 500 - - STC安全扫描: - stage: 前置检查 - component: stc - inputs: - tenantName: null - config: - execute: - isAllowSkip: true - - 构建并发布到测试库: - stage: 构建&发布到测试库 - id: build - component: pypi-artifact-uploader - inputs: - buildImage: reg.docker.alibaba-inc.com/aii/aistudio:aistudio-190677225-3221750112-1752554942251 - # buildImage: reg.docker.alibaba-inc.com/aii/aistudio:12150173-20251107143737 # max v2 - # buildTool: poetry - # artifactType: "wheel" # 仅打wheel包,如需要 tgz,请删除此行 - registry: "https://artifacts.antgroup-inc.cn/artifact/repositories/simple-dev/" # 测试库地址 - # workdir: . # pypi 工程目录,默认在此目录下进行 python -m build 并输出到 dist 目录, 详细可参考组件首页说明 - buildCmd: "python setup.py bdist_wheel" - only: - - master - - 选择迁移制品: - id: check - stage: 发布正式库前检查 - component: artifact-transfer-check - inputs: - artifactsConfigs: - - artifacts: ${{jobs.build.outputs.artifacts}} - antArtifactRepo: simple - - 同步至正式库: - stage: 发布正式库 - component: ant-artifact-transfer - inputs: - transferArtifacts: ${{jobs.check.outputs.transferArtifacts}} - config: - beforeExecute: - isAutoSkip: ${{jobs.check.outputs.selectedCount}} = 0 - confirm: - buttonName: 确认发布正式库 - approvers: ["nanxiao.zy", "liangchen.liangche"] # 具备发布角色的用户的域账号, 如 ["yunjie.gyj"] From 9e0ea1c791753fe355fb6c93f94cef4c5779c0e1 Mon Sep 17 00:00:00 2001 From: Yao Zhao Date: Fri, 23 Jan 2026 15:15:21 +0800 Subject: [PATCH 05/11] Update README.md --- README.md | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/README.md b/README.md index cdd7d39..239f7ea 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 Zhen Wang and Longfei Li} +} +``` From d65c3e5cea9c51c53472f857d3f088f69d60a279 Mon Sep 17 00:00:00 2001 From: Yao Zhao Date: Fri, 23 Jan 2026 15:16:35 +0800 Subject: [PATCH 06/11] Update README.md --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 239f7ea..86ad03a 100644 --- a/README.md +++ b/README.md @@ -74,6 +74,6 @@ Please refer to [API](https://inclusionai.github.io/linghe/) ``` @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 Zhen Wang and Longfei Li} +author={Yao Zhao and Chen Liang and Jingyu Hu and Zixuan Cheng and Longfei Li} } ``` From 02edabd91feb85b02e8d62ad690a15138efe6cbe Mon Sep 17 00:00:00 2001 From: long0x0 Date: Wed, 11 Feb 2026 22:07:35 +0800 Subject: [PATCH 07/11] format the code. --- benchmark/bench_gemm.py | 217 +- benchmark/bench_grad_norm.py | 43 +- benchmark/bench_la.py | 30 +- benchmark/bench_loss.py | 121 +- benchmark/bench_mla.py | 309 +- benchmark/bench_mla_rope.py | 327 +- benchmark/bench_norm.py | 109 +- benchmark/bench_permutation.py | 221 +- benchmark/bench_quantization.py | 84 +- benchmark/bench_topk.py | 86 +- linghe/attn/la.py | 792 ++-- linghe/attn/mla.py | 1266 ++++--- linghe/experimental/demb.py | 270 +- linghe/experimental/dla.py | 529 +-- linghe/experimental/dmm.py | 113 +- .../experimental/gmem_barrier_arrive_wait.py | 29 +- linghe/experimental/symm_mem_barrier.py | 24 +- linghe/experimental/test_demb.py | 267 +- linghe/experimental/test_dla.py | 119 +- linghe/experimental/test_dmm.py | 88 +- linghe/facade/emb.py | 21 +- linghe/facade/fp32_gemm.py | 11 +- linghe/facade/gate.py | 34 +- linghe/facade/hadamard_quant_linear.py | 110 +- linghe/facade/loss.py | 80 +- linghe/facade/mla.py | 157 +- linghe/facade/norm.py | 9 +- linghe/facade/permutation.py | 109 +- linghe/facade/rope.py | 319 +- linghe/facade/silu.py | 108 +- linghe/facade/smooth_quant_linear.py | 120 +- linghe/facade/topk.py | 74 +- linghe/facade/transpose.py | 2 +- linghe/gemm/blockwise_fp8_gemm.py | 142 +- linghe/gemm/channelwise_fp8_gemm.py | 74 +- linghe/gemm/fp32_gemm.py | 349 +- linghe/quant/block.py | 222 +- linghe/quant/channel.py | 111 +- linghe/quant/group.py | 17 +- linghe/quant/hadamard.py | 100 +- linghe/quant/smooth.py | 695 ++-- linghe/tools/benchmark.py | 59 +- linghe/tools/check.py | 74 +- linghe/tools/util.py | 220 +- linghe/utils/add.py | 68 +- linghe/utils/emb.py | 209 +- linghe/utils/gate.py | 129 +- linghe/utils/gather.py | 580 +-- linghe/utils/loss.py | 318 +- linghe/utils/mul.py | 52 +- linghe/utils/norm.py | 520 +-- linghe/utils/rearange.py | 34 +- linghe/utils/reduce.py | 129 +- linghe/utils/rope.py | 3173 ++++++++++------- linghe/utils/scatter.py | 103 +- linghe/utils/silu.py | 1482 ++++---- linghe/utils/topk.py | 110 +- linghe/utils/transpose.py | 250 +- linghe/utils/unary.py | 54 +- scripts/plot_input_output.py | 128 +- scripts/reproduce_triton_bug.py | 80 +- setup.py | 5 +- tests/test_add.py | 42 +- tests/test_blockwise_fp8_gemm.py | 57 +- tests/test_blockwise_quant.py | 106 +- tests/test_channel_quant.py | 50 +- tests/test_channelwise_fp8_gemm.py | 129 +- tests/test_dist_loss.py | 215 +- tests/test_embedding.py | 156 +- tests/test_fp32_gemm.py | 170 +- tests/test_gate.py | 216 +- tests/test_gather.py | 780 ++-- tests/test_group_quant.py | 13 +- tests/test_hadamard_quant.py | 58 +- tests/test_la.py | 108 +- tests/test_loss.py | 244 +- tests/test_mla.py | 548 ++- tests/test_mul.py | 73 +- tests/test_norm.py | 316 +- tests/test_rearange.py | 51 +- tests/test_reduce.py | 90 +- tests/test_rope.py | 1042 ++++-- tests/test_scatter.py | 43 +- tests/test_silu.py | 780 ++-- tests/test_smooth_quant.py | 423 ++- tests/test_topk.py | 283 +- tests/test_transpose.py | 129 +- tests/test_unary.py | 101 +- 88 files changed, 12585 insertions(+), 9323 deletions(-) diff --git a/benchmark/bench_gemm.py b/benchmark/bench_gemm.py index 17ff702..cb624a0 100644 --- a/benchmark/bench_gemm.py +++ b/benchmark/bench_gemm.py @@ -17,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 @@ -30,7 +30,7 @@ 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 @@ -38,7 +38,7 @@ def torch_accum_weight(x, w, out, x_scale, w_scale): 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) @@ -54,79 +54,111 @@ def bench_cublas_channelwise_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 bench_te_blockwise_gemm(M=4096, N=4096, K=4096, round_scale=False): # layout == 'TN': # forward, y=x@w 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, Float8BlockQuantizer + 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 - quantizer = Float8BlockQuantizer(TE_DType[torch.float8_e4m3fn], - rowwise=True, - columnwise=True, amax_epsilon=0, - force_pow_2_scales=round_scale, - block_scaling_dim=1) + 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' + 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) + 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 - ) + 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 - ) + 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.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.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=}') @@ -136,7 +168,7 @@ def bench_te_blockwise_gemm(M=4096, N=4096, K=4096, round_scale=False): transa = True B = qx transb = False - # out = torch.randn( (M, N), device='cuda:0', dtype=torch.bfloat16) + # out = torch.randn( (M, N), device='cuda:0', dtype=torch.bfloat16) out = None quantization_params = None out_dtype = TE_DType[torch.bfloat16] @@ -179,18 +211,20 @@ def bench_te_blockwise_gemm(M=4096, N=4096, K=4096, round_scale=False): ref_out = x @ w.t() - rel_err = ( - out - ref_out).abs().sum().item() / ref_out.abs().sum().item() + 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}') + 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) + 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): @@ -204,31 +238,33 @@ def bench_te_mxfp8_gemm(M=4096, N=4096, K=4096): 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 = 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) + 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, - ) + 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 @@ -269,11 +305,12 @@ def bench_te_mxfp8_gemm(M=4096, N=4096, K=4096): 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) + benchmark_func( + tex.generic_gemm, *args, n_repeat=100, ref_flops=ref_flops, ref_bytes=ref_bytes + ) -if __name__ == '__main__': +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) diff --git a/benchmark/bench_grad_norm.py b/benchmark/bench_grad_norm.py index d3611e0..2b47cc4 100644 --- a/benchmark/bench_grad_norm.py +++ b/benchmark/bench_grad_norm.py @@ -1,8 +1,10 @@ import random import torch -from transformer_engine.pytorch.optimizers import multi_tensor_applier, \ - multi_tensor_l2norm +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 @@ -10,10 +12,12 @@ 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)] + 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') + dummy_overflow_buf = torch.tensor([0], dtype=torch.int, device="cuda") grad_norm_ref, _ = multi_tensor_applier( multi_tensor_l2norm, dummy_overflow_buf, @@ -22,20 +26,23 @@ def bench_batch_norm(M=4096, N=2048, k=32): ) grad_norm = triton_batch_norm(xs, ord=2, norm=True) - output_check(grad_norm_ref[0], grad_norm, 'l2_norm') + 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__': + 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 index c2cb67a..d991d01 100644 --- a/benchmark/bench_la.py +++ b/benchmark/bench_la.py @@ -5,13 +5,16 @@ 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) + 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( @@ -26,10 +29,17 @@ def bench_la(B=1, S=4096, H=32, D=128): head_first=False, ) - benchmark_func(chunk_lightning_attn, query, key, value, layer_idx=1, - num_layers=20, output_final_state=True) + 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__': +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 e6047cc..2d45ac1 100644 --- a/benchmark/bench_loss.py +++ b/benchmark/bench_loss.py @@ -7,34 +7,31 @@ 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.check import output_check -from linghe.utils.loss import (triton_softmax_cross_entropy_backward, - triton_softmax_cross_entropy_forward) +from linghe.utils.loss import ( + triton_softmax_cross_entropy_backward, + triton_softmax_cross_entropy_forward, +) def fused_cross_entropy_forward_backward(logits, targets, input_grad, pg): - losses = fused_vocab_parallel_cross_entropy(logits[None], - targets[None], - pg)[0] + losses = fused_vocab_parallel_cross_entropy(logits[None], targets[None], pg)[0] losses.backward(input_grad) return losses, logits.grad def te_cross_entropy_forward_backward(logits, targets, input_grad): - losses = parallel_cross_entropy(logits[None], - targets[None]) + 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, - inplace=True): +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, @@ -48,60 +45,86 @@ def triton_cross_entropy_forward_backward(logits, targets, input_grad, def bench_triton_softmax_cross_entropy(M=4096, N=157184): - device = 'cuda:0' + 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.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') + pg = dist.new_group(ranks=[0], backend="nccl") fused_losses, fused_grad = fused_cross_entropy_forward_backward( - logits.detach().clone().requires_grad_(), targets, input_grad, pg) + 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) + 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) + logits.detach().clone().requires_grad_(), targets, input_grad, inplace=False + ) output_check(fused_losses, triton_losses) 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__': + 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 index 3063697..b1d5e71 100644 --- a/benchmark/bench_mla.py +++ b/benchmark/bench_mla.py @@ -14,28 +14,28 @@ 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: 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 @@ -55,8 +55,9 @@ def __init__( 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.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 @@ -69,14 +70,14 @@ def __init__( 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, + 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 @@ -86,44 +87,51 @@ def _run_dot_product_attention( 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" + 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" + 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" + 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_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" + 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" + [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" + [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 = 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) @@ -134,8 +142,9 @@ def _run_dot_product_attention( 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 + 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) @@ -151,8 +160,8 @@ def _run_dot_product_attention( [ attention_mask_q, torch.Tensor( - [False] * seqlens_q[i] + [True] * ( - config.max_seqlen_q - seqlens_q[i]) + [False] * seqlens_q[i] + + [True] * (config.max_seqlen_q - seqlens_q[i]) ) .to(dtype=torch.bool) .unsqueeze(0) @@ -170,8 +179,8 @@ def _run_dot_product_attention( [ attention_mask_q, torch.Tensor( - [False] * seqlens_q[i] + [True] * ( - config.max_seqlen_q - seqlens_q[i]) + [False] * seqlens_q[i] + + [True] * (config.max_seqlen_q - seqlens_q[i]) ) .to(dtype=torch.bool) .unsqueeze(0) @@ -247,13 +256,11 @@ def _run_dot_product_attention( 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[pad_range[0] : pad_range[1]] = 0.0 tensor_orig = torch.cat( - [tensor_orig, tensor[valid_range[0]: valid_range[1]]], - dim=0 + [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"]: + 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], @@ -263,10 +270,9 @@ def _run_dot_product_attention( 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[pad_range[0] : pad_range[1]] = 0.0 tensor_orig = torch.cat( - [tensor_orig, tensor[valid_range[0]: valid_range[1]]], - dim=0 + [tensor_orig, tensor[valid_range[0] : valid_range[1]]], dim=0 ) tensor_count = 1 split_dim = 0 @@ -275,11 +281,11 @@ def _run_dot_product_attention( tensor_count = int(l) split_dim = dim break - tensors = torch.split(tensor, 1, dim=split_dim) if split_dim != 0 else [ - tensor] + 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] + 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: @@ -297,10 +303,10 @@ def _run_dot_product_attention( 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_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) @@ -310,12 +316,13 @@ def _run_dot_product_attention( 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 + 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 + [out_grad_orig, out_grad[valid_range[0] : valid_range[1]]], dim=0 ) # Create bias @@ -374,8 +381,12 @@ def _run_dot_product_attention( 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, + 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, @@ -415,19 +426,20 @@ def _run_dot_product_attention( 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) + [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 + [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 + [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 + [v_grad_orig, v.grad[valid_range_kv[0] : valid_range_kv[1]]], + dim=0, ) if is_training: return ( @@ -439,8 +451,7 @@ def _run_dot_product_attention( 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) + return out, max_logit, (q.grad, k.grad, v.grad, d_softmax_offset) else: return out, max_logit, (None, None, None, d_softmax_offset) @@ -452,16 +463,16 @@ def fused_attn(block, q, k, v, cu_seqlens_q, cu_seqlens_kv, mask): v, window_size=(-1, 0), attention_mask=mask, - qkv_format='thd', + 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', + attn_mask_type="padding_causal", checkpoint_core_attention=False, - core_attention_bias_type='no_bias', + core_attention_bias_type="no_bias", core_attention_bias=None, alibi_slopes=None, fast_zero_fill=True, @@ -471,50 +482,65 @@ def fused_attn(block, q, k, v, cu_seqlens_q, cu_seqlens_kv, mask): 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), - ) + 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' + 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) + _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), - ) + 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' + qkv_layout = "thd_thd_thd" workspace_opt = True pad_between_seqs = True is_training = True @@ -525,58 +551,55 @@ def bench_fused_attn(B=1, S=8192, H=64): 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" + 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', + 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', + 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" + [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" + [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) + 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) + 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( + 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) + benchmark_func(out.backward, g, retain_graph=True, n_profile=1) -if __name__ == '__main__': +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 index a2aecf4..71ef061 100644 --- a/benchmark/bench_mla_rope.py +++ b/benchmark/bench_mla_rope.py @@ -10,20 +10,24 @@ 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 bench_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=True): dtype = torch.bfloat16 - device = 'cuda:0' + 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_() + 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] @@ -63,7 +67,7 @@ def bench_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=True): 128, cu_seqlens_kv=None, cp_rank=0, - cp_size=1 + cp_size=1, ) if transpose: query_ref = query_ref.transpose(0, 1) @@ -77,16 +81,18 @@ def bench_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=True): 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, 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) @@ -94,64 +100,87 @@ def bench_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=True): 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(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) + 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) + 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): +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_() + 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_() + 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_() + 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) + 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 = 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) + 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 @@ -190,16 +219,18 @@ def bench_varlen_mla_rope(lengths=[2048, 2048], H=32, rope_theta=10000.0, 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, 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) @@ -208,65 +239,117 @@ def bench_varlen_mla_rope(lengths=[2048, 2048], H=32, rope_theta=10000.0, 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(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) + 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( + 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) + 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__': +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) + 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 index 45042cb..ee1c91a 100644 --- a/benchmark/bench_norm.py +++ b/benchmark/bench_norm.py @@ -1,10 +1,11 @@ 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 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 @@ -17,7 +18,7 @@ def bench_rmsnorm(bs=1, M=4096, N=4096): # M, N, K = 4096, 8192, 4096 dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 x = torch.randn(bs, M, N, dtype=dtype, requires_grad=True, device=device) @@ -25,10 +26,7 @@ def bench_rmsnorm(bs=1, M=4096, N=4096): 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' + normalized_shape=N, eps=1e-6, dtype=torch.bfloat16, device="cuda" ) rmsnorm_torch = torch.compile(rmsnorm_torch) @@ -50,43 +48,78 @@ def triton_forward_backward(x_triton_back, g_triton_back, dy): 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) + 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) + 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 = 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) + 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(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) + benchmark_func( + triton_forward_backward, x, weight, dy, n_repeat=n_repeat, ref_time=ref_time + ) -if __name__ == '__main__': +if __name__ == "__main__": bench_rmsnorm(1, 4096, 4096) diff --git a/benchmark/bench_permutation.py b/benchmark/bench_permutation.py index 7b86295..03b146b 100644 --- a/benchmark/bench_permutation.py +++ b/benchmark/bench_permutation.py @@ -2,18 +2,22 @@ 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 transformer_engine.pytorch.tensor.float8_blockwise_tensor import ( + Float8BlockQuantizer, +) from linghe.tools.benchmark import benchmark_func from linghe.tools.util import torch_make_indices -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, - ) +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): @@ -33,50 +37,57 @@ 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): +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) + 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) + 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) +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) scales = torch.randn(M, dtype=dtype, device=device) @@ -84,59 +95,94 @@ def bench_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.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) + 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' + device = "cuda:0" dtype = torch.bfloat16 x = torch.randn(M, N, dtype=dtype, device=device) fp8_padding = Fp8Padding(32, 16) 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) - 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) + 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) + 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) @@ -148,22 +194,37 @@ def bench_triton_unpermute_with_mask_map(M=4098, N=4096, n_experts=32, topk=2): 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__': + 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_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 index 3515963..79f75fd 100644 --- a/benchmark/bench_quantization.py +++ b/benchmark/bench_quantization.py @@ -3,8 +3,9 @@ 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.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 @@ -14,18 +15,20 @@ 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) + 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' + 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.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 @@ -33,56 +36,55 @@ def bench_blockwise_quantization(M=8192, N=4096, round_scale=True): 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') + 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) + 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.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') + 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' + device = "cuda:0" - splits = [max(random.randint(M - 256, M + 256), 0) for x in - range(n_experts)] + 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)) + 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) + 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') @@ -92,11 +94,17 @@ def bench_batch_mxfp8_quant(M=4096, N=4096, n_experts=32, bench=False): 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) + benchmark_func( + triton_batch_mxfp8_quant, + x, + token_count_per_expert, + splits, + output_mode=2, + ref_bytes=ref_bytes, + ) -if __name__ == '__main__': +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) diff --git a/benchmark/bench_topk.py b/benchmark/bench_topk.py index 377c7cd..4204653 100644 --- a/benchmark/bench_topk.py +++ b/benchmark/bench_topk.py @@ -11,31 +11,21 @@ def bench_topk(M=4096, N=256, k=8): - device = 'cuda:0' + 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) + values_ref, indices_ref = torch.topk(logits, k) # print(f'{indices_ref[0]}') - values, indices = fused_topk(logits, - k) + 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) + 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' + 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_() @@ -45,7 +35,7 @@ def bench_group_topk_score(M=4096, N=256, k=8): group_topk = 4 scaling_factor = 2.5 deterministic_mode = True - score_function = 'sigmoid' + score_function = "sigmoid" expert_bias = torch.randn((N,), dtype=torch.float32, device=device) moe_router_fusion = True @@ -62,7 +52,8 @@ def bench_group_topk_score(M=4096, N=256, k=8): deterministic_mode, score_function, expert_bias, - moe_router_fusion) + moe_router_fusion, + ) # print((-routing_map_ref[0].float()).argsort(0)) probs, routing_map, tokens_per_expert = group_topk_score( @@ -72,37 +63,42 @@ def bench_group_topk_score(M=4096, N=256, k=8): num_groups=num_groups, group_topk=group_topk, scaling_factor=scaling_factor, - score_function=score_function) + 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) + 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) + 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__': +if __name__ == "__main__": bench_topk(M=8192, N=256, k=8) bench_group_topk_score(M=8192, N=256, k=8) diff --git a/linghe/attn/la.py b/linghe/attn/la.py index 0d6af35..2f032ba 100644 --- a/linghe/attn/la.py +++ b/linghe/attn/la.py @@ -10,22 +10,22 @@ @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, + 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) @@ -44,40 +44,40 @@ def fp32_lightning_attention_forward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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) @@ -105,30 +105,31 @@ def fp32_lightning_attention_forward_kernel( 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.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, + 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) @@ -147,46 +148,45 @@ def lightning_attention_forward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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 + 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 @@ -209,8 +209,9 @@ def lightning_attention_forward_kernel( 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.atomic_add( + out_ptrs + n * H * D, o.to(Out.dtype.element_ty), sem="relaxed" + ) tl.store(s_ptrs, state) @@ -228,8 +229,9 @@ def _output_sum_kernel(T, O, DIM: tl.constexpr, NUM_BLOCK: tl.constexpr): 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): +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" @@ -249,20 +251,20 @@ def triton_lightning_attention_forward(q, k, v, decay_scales, hpc=False, 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 - ) + 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 - ) + 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 + 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, @@ -295,22 +297,22 @@ def triton_lightning_attention_forward(q, k, v, decay_scales, hpc=False, @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, + 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) @@ -329,39 +331,39 @@ def fp32_lightning_attention_q_backward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + DQ + + c0 * D * H + + hid * D + + kid * KD + + (offs_b[:, None] * H * D + offs_k[None, :]) ) state = tl.zeros((KD, VD), dtype=tl.float32) @@ -390,36 +392,39 @@ def fp32_lightning_attention_q_backward_kernel( 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 + 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') + 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, + 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) @@ -438,38 +443,37 @@ def lightning_attention_q_backward_kernel( offs_v = tl.arange(0, VD) k_ptrs = ( - K - + c0 * stride_k - + hid * D - + kid * KD - + (offs_b[:, None] * stride_k + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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 + mask = tl.where(offs_b[None, :] <= offs_b[:, None], mask, 0.0) * softmax_scale decay_offs = BLOCK - 1 - offs_b @@ -493,37 +497,40 @@ def lightning_attention_q_backward_kernel( 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 + 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') + 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, + 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) @@ -542,47 +549,47 @@ def fp32_lightning_attention_kv_backward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + DV + + c0 * H * D + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) ) gs = tl.zeros((KD, VD), dtype=tl.float32) @@ -613,8 +620,7 @@ def fp32_lightning_attention_kv_backward_kernel( 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 + 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] @@ -625,34 +631,36 @@ def fp32_lightning_attention_kv_backward_kernel( 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') + 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.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, + 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) @@ -671,47 +679,47 @@ def lightning_attention_kv_backward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + DV + + c0 * H * D + + hid * D + + vid * VD + + (offs_b[:, None] * H * D + offs_v[None, :]) ) gs = tl.zeros((KD, VD), dtype=tl.float32) @@ -725,8 +733,7 @@ def lightning_attention_kv_backward_kernel( 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 + 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): @@ -757,18 +764,20 @@ def lightning_attention_kv_backward_kernel( 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') + 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.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): +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) @@ -786,14 +795,16 @@ def triton_lightning_attention_backward(output_grad, q, k, v, decay_scales, (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 - ) + 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 = ( + fp32_lightning_attention_q_backward_kernel + if hp + else lightning_attention_q_backward_kernel + ) kernel[grid]( q, k, @@ -825,21 +836,21 @@ def triton_lightning_attention_backward(output_grad, q, k, v, decay_scales, (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 - ) + 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 - ) + 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 = ( + fp32_lightning_attention_kv_backward_kernel + if hp + else lightning_attention_kv_backward_kernel + ) kernel[grid]( q, k, @@ -872,25 +883,25 @@ def triton_lightning_attention_backward(output_grad, q, k, v, decay_scales, @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, + 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) @@ -909,61 +920,61 @@ def fused_lightning_attention_backward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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) @@ -1001,16 +1012,15 @@ def fused_lightning_attention_backward_kernel( 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 + 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 = ( + tl.dot(dqk, ks) + + tl.dot(g, tl.trans(state)) * softmax_scale * decays[:, None] + ) dq *= amps[:, None] state /= block_decay @@ -1024,20 +1034,23 @@ def fused_lightning_attention_backward_kernel( 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') + 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') + 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): +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) @@ -1058,20 +1071,14 @@ def triton_fused_lightning_attention_backward(output_grad, q, k, v, s, (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 - ) + 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 - ) + 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) @@ -1109,6 +1116,7 @@ def triton_fused_lightning_attention_backward(output_grad, q, k, v, s, dv = dv.to(dtype) return dq, dk, dv + # @triton.jit # def varlen_lightning_attention_forward_kernel( # Q, @@ -1284,7 +1292,7 @@ def triton_fused_lightning_attention_backward(output_grad, q, k, v, s, # ) # # BLOCK should <= 64 -# BLOCK = 32 +# 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]( @@ -1558,7 +1566,7 @@ def triton_fused_lightning_attention_backward(output_grad, q, k, v, s, # ) # # BLOCK should <= 64 -# BLOCK = 32 +# 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]( diff --git a/linghe/attn/mla.py b/linghe/attn/mla.py index 6534f02..18382dc 100644 --- a/linghe/attn/mla.py +++ b/linghe/attn/mla.py @@ -12,20 +12,20 @@ @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, + 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) @@ -40,39 +40,39 @@ def deprecated_mla_forward_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + (bid * L + mid * M) * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_0[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) ) q0 = tl.load(q0_ptrs) @@ -96,8 +96,7 @@ def deprecated_mla_forward_kernel( qk = tl.dot(q0, tl.trans(k0), qk) - qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], - 0.0, -1e9) + 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) @@ -110,10 +109,10 @@ def deprecated_mla_forward_kernel( # [B, L, H, 128] out_ptrs = ( - Out - + (bid * L + mid * M) * H * 128 - + hid * 128 - + (offs_m[:, None] * 128 * H + offs_0[None, :]) + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) ) tl.store(out_ptrs, acc_o) @@ -122,23 +121,23 @@ def deprecated_mla_forward_kernel( @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, + 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) @@ -153,24 +152,24 @@ def mla_forward_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + (bid * L + mid * M) * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_1[None, :]) + 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, :]) + 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, :]) + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) ) q0 = tl.load(q0_ptrs) @@ -202,8 +201,9 @@ def mla_forward_kernel( 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 += tl.where( + (mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], 0.0, -1e9 + ) qk *= softmax_scale @@ -231,17 +231,16 @@ def mla_forward_kernel( # [B, L, H, 128] out_ptrs = ( - Out - + (bid * L + mid * M) * H * 128 - + hid * 128 - + (offs_m[:, None] * 128 * H + offs_0[None, :]) + 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) + 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): @@ -299,20 +298,20 @@ def triton_mla_forward(q, k, v, causal=True, safe=True, clip_value=None): # 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, + 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) @@ -327,31 +326,31 @@ def naive_mla_ds_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + (bid * L + mid * M) * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_1[None, :]) + 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, :]) + 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, :]) + 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, :]) + GO + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) ) ds = tl.zeros((M,), dtype=tl.float32) @@ -381,8 +380,7 @@ def naive_mla_ds_kernel( qk = tl.dot(q2, tl.trans(k2), qk) - qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], - 0.0, -1e9) + 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) @@ -398,11 +396,11 @@ def naive_mla_ds_kernel( # dp and p dot sum @triton.jit def mla_ds_kernel( - G, - O, - DS, - L, - M: tl.constexpr, + G, + O, + DS, + L, + M: tl.constexpr, ): bid = tl.program_id(0) hid = tl.program_id(1) @@ -413,40 +411,40 @@ def mla_ds_kernel( 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, :]) - ) + 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) + 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, + 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) @@ -461,17 +459,17 @@ def deprecated_mla_backward_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + bid * L * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_1[None, :]) + 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, :]) + K + + (bid * L + nid * N) * stride_k + + hid * 192 + + (offs_n[:, None] * stride_k + offs_1[None, :]) ) k0 = tl.load(k0_ptrs) @@ -479,27 +477,27 @@ def deprecated_mla_backward_kernel( 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, :]) + 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, :]) + 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, :]) + 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) @@ -529,8 +527,9 @@ def deprecated_mla_backward_kernel( 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) + 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] @@ -555,20 +554,20 @@ def deprecated_mla_backward_kernel( 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, :]) + 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, :]) + 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) @@ -577,28 +576,28 @@ def deprecated_mla_backward_kernel( @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, + 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) @@ -614,80 +613,80 @@ def mla_backward_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + bid * L * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_0[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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 + + (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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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) @@ -704,8 +703,7 @@ def mla_backward_kernel( 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)) + 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) @@ -713,8 +711,9 @@ def mla_backward_kernel( 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.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: @@ -745,8 +744,8 @@ def mla_backward_kernel( 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') + 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) @@ -756,28 +755,28 @@ def mla_backward_kernel( 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, :]) + 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, :]) + 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, :]) + GK + + (bid * L + nid * N) * H * 192 + + hid * 192 + + 128 + + (offs_n[:, None] * 192 * H + offs_1[None, :]) ) tl.store(gk1_ptrs, dk1) @@ -785,12 +784,7 @@ def mla_backward_kernel( # ragged sum @triton.jit def mla_rs_kernel( - Q, - O, - H: tl.constexpr, - N: tl.constexpr, - BLOCK: tl.constexpr, - CAUSAL: tl.constexpr + 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) @@ -801,13 +795,7 @@ def mla_rs_kernel( 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 - ) + 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) @@ -817,21 +805,27 @@ def mla_rs_kernel( 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 - ) + 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): +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] @@ -852,24 +846,18 @@ def triton_mla_backward(go, o, q, k, v, lse, max_logits, causal=True, safe=True, 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 - ) + 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) + 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) + 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) @@ -917,41 +905,42 @@ def triton_mla_backward(go, o, q, k, v, lse, max_logits, causal=True, safe=True, 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, - ) + 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, + 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) @@ -982,24 +971,18 @@ def varlen_mla_forward_kernel( # [T, H, 192】 q0_ptrs = ( - Q - + (c0 + mid * M) * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_1[None, :]) + 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, :]) + 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, :]) + V + c0 * stride_v + hid * 128 + (offs_n[:, None] * stride_v + offs_0[None, :]) ) m_mask = (mid * M + offs_m) < length @@ -1035,8 +1018,11 @@ def varlen_mla_forward_kernel( if CAUSAL: qk += tl.where( - ((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :]) & ( - n_mask[None, :]), 0.0, -1e9) + ((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) @@ -1066,23 +1052,30 @@ def varlen_mla_forward_kernel( # [T, H, 128] out_ptrs = ( - Out - + (c0 + mid * M) * H * 128 - + hid * 128 - + (offs_m[:, None] * H * 128 + offs_0[None, :]) + 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): + 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] @@ -1137,31 +1130,31 @@ def triton_varlen_mla_forward(q, k, v, cu_seqlens, max_q_length, @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, + 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) @@ -1194,81 +1187,75 @@ def varlen_mla_backward_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + c0 * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_0[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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 + + (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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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) @@ -1286,11 +1273,11 @@ def varlen_mla_backward_kernel( 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) + 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) + 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]) @@ -1299,8 +1286,11 @@ def varlen_mla_backward_kernel( if CAUSAL: qk = tl.where( - ((m + offs_m)[:, None] >= (nid * N + offs_n)[None, :]) & ( - n_mask[None, :]), 0.0, -1e9) + ((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: @@ -1332,10 +1322,12 @@ def varlen_mla_backward_kernel( 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') + 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]) @@ -1345,27 +1337,27 @@ def varlen_mla_backward_kernel( 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, :]) + 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, :]) + 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, :]) + 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]) @@ -1373,15 +1365,15 @@ def varlen_mla_backward_kernel( # 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 + 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) @@ -1389,19 +1381,14 @@ def varlen_mla_rs_kernel( 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) + c1 = tl.min(tl.where(cu <= c0, 2**24, cu), 0) 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 - ) + 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) @@ -1412,24 +1399,30 @@ def varlen_mla_rs_kernel( 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 - ) + 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 - ): +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] @@ -1451,25 +1444,21 @@ def triton_varlen_mla_backward(go, o, q, k, v, lse, 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 - ) + 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) + 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) + 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) @@ -1518,41 +1507,42 @@ def triton_varlen_mla_backward(go, o, q, k, v, lse, num_warps = 4 num_stages = 3 PB = max(triton.next_power_of_2(B), 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, - ) + 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, + 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) @@ -1567,56 +1557,45 @@ def deprecated_fp8_mla_forward_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + (bid * L + mid * M) * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_0[None, :]) + 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, :]) + 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 - ) + 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, :]) + 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, :]) + 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 - ) + 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, :]) + V + + bid * L * stride_v + + hid * 128 + + (offs_n[:, None] * stride_v + offs_0[None, :]) ) q0 = tl.load(q0_ptrs) @@ -1642,8 +1621,7 @@ def deprecated_fp8_mla_forward_kernel( qk = tl.dot(q0, tl.trans(k0), qk) - qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], - 0.0, -1e9) + 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) @@ -1657,10 +1635,10 @@ def deprecated_fp8_mla_forward_kernel( # [B, L, H, 128] out_ptrs = ( - Out - + (bid * L + mid * M) * H * 128 - + hid * 128 - + (offs_m[:, None] * 128 * H + offs_0[None, :]) + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_0[None, :]) ) tl.store(out_ptrs, acc_o) @@ -1669,23 +1647,23 @@ def deprecated_fp8_mla_forward_kernel( @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, + 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) @@ -1700,41 +1678,30 @@ def padding_fp8_mla_forward_kernel( # [B, L, H, 192】 q0_ptrs = ( - Q - + (bid * L + mid * M) * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_0[None, :]) + 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 - ) + 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, :]) + 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 - ) + 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, :]) + 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, :]) @@ -1755,8 +1722,7 @@ def padding_fp8_mla_forward_kernel( qk = tl.dot(q0, tl.trans(k0)) - qk += tl.where((mid * M + offs_m)[:, None] >= (n + offs_n)[None, :], - 0.0, -1e9) + 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) @@ -1771,10 +1737,10 @@ def padding_fp8_mla_forward_kernel( # [B, L, H, 128] out_ptrs = ( - Out - + (bid * L + mid * M) * H * 128 - + hid * 128 - + (offs_m[:, None] * 128 * H + offs_1[None, :]) + Out + + (bid * L + mid * M) * H * 128 + + hid * 128 + + (offs_m[:, None] * 128 * H + offs_1[None, :]) ) tl.store(out_ptrs, acc_o) @@ -1783,23 +1749,23 @@ def padding_fp8_mla_forward_kernel( @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, + 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) @@ -1814,51 +1780,35 @@ def fp8_mla_forward_kernel( # [B, L, H, 192] q0_ptrs = ( - Q - + (bid * L + mid * M) * stride_q - + hid * 192 - + (offs_m[:, None] * stride_q + offs_1[None, :]) + 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 - ) + 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, :]) + 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 - ) + 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, :]) + 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 - ) + vs_ptrs = VS + bid * H * L + hid * L + offs_n q0 = tl.load(q0_ptrs) q1 = tl.load(q0_ptrs + 64) @@ -1883,8 +1833,9 @@ def fp8_mla_forward_kernel( 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.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) @@ -1913,16 +1864,17 @@ def fp8_mla_forward_kernel( 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, :]) + 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): +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] diff --git a/linghe/experimental/demb.py b/linghe/experimental/demb.py index b86edf4..9d97108 100644 --- a/linghe/experimental/demb.py +++ b/linghe/experimental/demb.py @@ -17,16 +17,18 @@ @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): +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)) @@ -35,8 +37,7 @@ def tp_embedding_lookup_forward_kernel(input_ids_ptr, 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.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) @@ -60,7 +61,8 @@ def tp_embedding_lookup_forward_kernel(input_ids_ptr, hasSubsequentMemAccess=True, ) buffer_ptr = tl.load(buffer_ptrs + input_id // V).to( - tl.pointer_type(tl.bfloat16)) + 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) @@ -84,8 +86,7 @@ def triton_tp_embedding_lookup_forward(input_ids, weights, hdl, group): if len(shape) == 2: M = shape[0] * shape[1] - outputs = torch.empty((shape[0], shape[1], d), device=device, - dtype=dtype) + 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) @@ -110,24 +111,25 @@ def triton_tp_embedding_lookup_forward(input_ids, weights, hdl, group): @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 - ): +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)) @@ -153,9 +155,9 @@ def tp_embedding_lookup_backward_kernel(grad_output_ptr, 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) + 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, @@ -168,14 +170,11 @@ def tp_embedding_lookup_backward_kernel(grad_output_ptr, for j in range(SIZE): if j != RANK: - buffer_ptr = tl.load(buffer_ptrs + j).to( - tl.pointer_type(tl.bfloat16)) + 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) + 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) + tl.store(grad_ptr + input_id % V * DIM + tl.arange(0, DIM), outputs, mask=mask) else: @@ -184,13 +183,12 @@ def tp_embedding_lookup_backward_kernel(grad_output_ptr, 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) + 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.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( @@ -203,8 +201,9 @@ def tp_embedding_lookup_backward_kernel(grad_output_ptr, ) -def triton_tp_embedding_lookup_backward(grad_output, x, g_ptr, vocab_size, hdl, - group, dtype=torch.bfloat16): +def triton_tp_embedding_lookup_backward( + grad_output, x, g_ptr, vocab_size, hdl, group, dtype=torch.bfloat16 +): """ inplace update embedding weight gradient Args: @@ -251,22 +250,24 @@ def triton_tp_embedding_lookup_backward(grad_output, x, g_ptr, vocab_size, hdl, group_size, group_rank, num_stages=num_stages, - num_warps=num_warps + 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): +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 @@ -275,8 +276,7 @@ def sp_embedding_lookup_forward_kernel(input_ids_ptr, 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) + 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( @@ -288,17 +288,17 @@ def sp_embedding_lookup_forward_kernel(input_ids_ptr, hasSubsequentMemAccess=True, ) buffer_ptr = tl.load(buffer_ptrs + input_id // V).to( - tl.pointer_type(tl.bfloat16)) + 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) + 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) + 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)) + 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( @@ -329,15 +329,16 @@ def triton_sp_embedding_lookup_forward(input_ids, weights, hdl, group): 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) + 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) + 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 @@ -362,32 +363,32 @@ def triton_sp_embedding_lookup_forward(input_ids, weights, hdl, group): @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 - ): +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)) + c01 = tl.load(accum_counts_ptr + chunk * (L + 1) + pid + tl.arange(0, 2)) c0, c1 = tl.split(c01) if c0 != c1: @@ -408,13 +409,20 @@ def sp_embedding_lookup_backward_kernel(grad_output_ptr, 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) + 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') + outputs, + mask=mask, + sem="relaxed", + ) else: @@ -423,16 +431,22 @@ def sp_embedding_lookup_backward_kernel(grad_output_ptr, 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) + 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)) + 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) + tl.store( + buffer_ptr + pid * DIM + tl.arange(0, DIM), outputs, mask=mask + ) symm_mem_sync( signal_ptrs, None, @@ -453,19 +467,28 @@ def sp_embedding_lookup_backward_kernel(grad_output_ptr, ) buffer_ptr = tl.load(buffer_ptrs + RANK).to( - tl.pointer_type(tl.float32)) + 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) + 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') + 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): +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: @@ -489,23 +512,24 @@ def triton_sp_embedding_lookup_backward(grad_output, input_ids, g_ptr, 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) + 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) + 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) + 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) @@ -532,5 +556,5 @@ def triton_sp_embedding_lookup_backward(grad_output, input_ids, g_ptr, group_size, group_rank, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) diff --git a/linghe/experimental/dla.py b/linghe/experimental/dla.py index cd45da7..5fcea16 100644 --- a/linghe/experimental/dla.py +++ b/linghe/experimental/dla.py @@ -16,26 +16,26 @@ @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, + 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) @@ -56,54 +56,53 @@ def cp_lightning_attention_forward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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 + 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 @@ -131,7 +130,7 @@ def cp_lightning_attention_forward_kernel( if KD == D: tl.store(out_ptrs + n * H * D, o) else: - tl.atomic_add(out_ptrs + n * H * D, o, sem='relaxed') + 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) @@ -160,7 +159,7 @@ def cp_lightning_attention_forward_kernel( if KD == D: tl.store(out_ptrs + n * H * D, o) else: - tl.atomic_add(out_ptrs + n * H * D, o, sem='relaxed') + tl.atomic_add(out_ptrs + n * H * D, o, sem="relaxed") tl.store(buffer_ptr + DD + buffer_offs, state1) symm_mem_sync( @@ -177,8 +176,7 @@ def cp_lightning_attention_forward_kernel( 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)) + 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) @@ -187,8 +185,7 @@ def cp_lightning_attention_forward_kernel( 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)) + 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) @@ -207,7 +204,8 @@ def cp_lightning_attention_forward_kernel( 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)) + tl.pointer_type(tl.float32) + ) state0 += tl.load(pre_buffer_ptr + buffer_offs) * chunk_decay tl.store(buffer_ptr + buffer_offs, state0) @@ -216,7 +214,8 @@ def cp_lightning_attention_forward_kernel( 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)) + tl.pointer_type(tl.float32) + ) state1 += tl.load(pre_buffer_ptr + DD + buffer_offs) * chunk_decay tl.store(buffer_ptr + DD + buffer_offs, state1) @@ -235,7 +234,8 @@ def cp_lightning_attention_forward_kernel( 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)) + tl.pointer_type(tl.float32) + ) state0 += tl.load(pre_buffer_ptr + buffer_offs) * chunk_decay tl.store(buffer_ptr + buffer_offs, state0) @@ -258,9 +258,9 @@ def cp_lightning_attention_forward_kernel( 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.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( @@ -282,9 +282,9 @@ def cp_lightning_attention_forward_kernel( 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.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 @@ -294,9 +294,9 @@ def cp_lightning_attention_forward_kernel( 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.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( @@ -314,19 +314,17 @@ def cp_lightning_attention_forward_kernel( 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)) + 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): + 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)) + 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) @@ -349,12 +347,12 @@ def cp_lightning_attention_forward_kernel( 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.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) + buffer_ptr + pre_offs + buffer_offs + ) # .to(Q.dtype.element_ty) iter_amps = amps * block_decay for n in range(0, L // 2, BLOCK): @@ -368,7 +366,7 @@ def cp_lightning_attention_forward_kernel( o = tl.dot(q, pre_state) iter_amps *= block_decay - tl.atomic_add(out_ptrs + n * H * D, o, sem='relaxed') + tl.atomic_add(out_ptrs + n * H * D, o, sem="relaxed") gcid = 2 * SIZE - 1 - RANK pre_gcid = gcid - 1 @@ -377,8 +375,7 @@ def cp_lightning_attention_forward_kernel( 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) + 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): @@ -392,11 +389,12 @@ def cp_lightning_attention_forward_kernel( o = tl.dot(q, pre_state) iter_amps *= block_decay - tl.atomic_add(out_ptrs + n * H * D, o, sem='relaxed') + 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): +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" @@ -424,9 +422,7 @@ def triton_cp_lightning_attention_forward(q, k, v, decay_scales, hdl, group, (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 - ) + 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 @@ -464,28 +460,28 @@ def triton_cp_lightning_attention_forward(q, k, v, decay_scales, hdl, group, @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, + 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) @@ -506,48 +502,48 @@ def cp_lightning_attention_q_backward_kernel( offs_v = tl.arange(0, VD) k_ptrs = ( - K - + c0 * stride_k - + hid * D - + kid * KD - + (offs_b[:, None] * stride_k + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + bid * stride_s + + hid * 2 * D * D + + kid * D * KD + + vid * VD + + (offs_k[:, None] * D + offs_v[None, :]) ) # store state to buffer @@ -571,8 +567,7 @@ def cp_lightning_attention_q_backward_kernel( ) 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 + 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] @@ -584,11 +579,9 @@ def cp_lightning_attention_q_backward_kernel( 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.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) + 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) @@ -601,14 +594,17 @@ def cp_lightning_attention_q_backward_kernel( 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 + 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') + 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) @@ -629,8 +625,7 @@ def cp_lightning_attention_q_backward_kernel( 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) + 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) @@ -643,41 +638,44 @@ def cp_lightning_attention_q_backward_kernel( 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 + 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') + 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, + 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) @@ -698,54 +696,54 @@ def cp_lightning_attention_kv_backward_kernel( offs_v = tl.arange(0, VD) q_ptrs = ( - Q - + c0 * stride_q - + hid * D - + kid * KD - + (offs_b[:, None] * stride_q + offs_k[None, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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, :]) + 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 @@ -756,8 +754,7 @@ def cp_lightning_attention_kv_backward_kernel( 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 + 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) @@ -789,13 +786,15 @@ def cp_lightning_attention_kv_backward_kernel( 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') + 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.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) @@ -831,13 +830,15 @@ def cp_lightning_attention_kv_backward_kernel( 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') + 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.atomic_add( + dv_ptrs + n * H * D, dv.to(DV.dtype.element_ty), sem="relaxed" + ) tl.store(buffer_ptr + DD + buffer_offs, gs1) @@ -857,7 +858,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs0 += tl.load(next_buffer_ptr + buffer_offs) * chunk_decay tl.store(buffer_ptr + buffer_offs, gs0) @@ -867,7 +869,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay tl.store(buffer_ptr + DD + buffer_offs, gs1) @@ -888,7 +891,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs0 += tl.load(next_buffer_ptr + buffer_offs) * chunk_decay tl.store(buffer_ptr + buffer_offs, gs0) @@ -898,7 +902,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay tl.store(buffer_ptr + DD + buffer_offs, gs1) @@ -918,7 +923,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay tl.store(buffer_ptr + DD + buffer_offs, gs1) @@ -941,9 +947,9 @@ def cp_lightning_attention_kv_backward_kernel( 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.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( @@ -965,9 +971,9 @@ def cp_lightning_attention_kv_backward_kernel( 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.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 @@ -977,9 +983,9 @@ def cp_lightning_attention_kv_backward_kernel( 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.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( @@ -999,7 +1005,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs0 += tl.load(next_buffer_ptr + next_offs + buffer_offs) * chunk_decay tl.store(buffer_ptr + buffer_offs, gs0) @@ -1012,7 +1019,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs1 += tl.load(next_buffer_ptr + DD + buffer_offs) * chunk_decay tl.store(buffer_ptr + DD + buffer_offs, gs1) @@ -1030,8 +1038,7 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + 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) @@ -1048,10 +1055,8 @@ def cp_lightning_attention_kv_backward_kernel( 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') + 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 @@ -1060,7 +1065,8 @@ def cp_lightning_attention_kv_backward_kernel( 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)) + tl.pointer_type(tl.float32) + ) gs = tl.load(next_buffer_ptr + next_offs + buffer_offs) n_steps = tl.cdiv(L // 2, BLOCK) @@ -1077,16 +1083,27 @@ def cp_lightning_attention_kv_backward_kernel( 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): + 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) diff --git a/linghe/experimental/dmm.py b/linghe/experimental/dmm.py index 20381f8..c700aab 100644 --- a/linghe/experimental/dmm.py +++ b/linghe/experimental/dmm.py @@ -13,21 +13,21 @@ @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, + 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) @@ -36,13 +36,11 @@ def split_tp_mm_kernel( 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_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] + 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): @@ -71,11 +69,9 @@ def split_tp_mm_kernel( 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.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, :]) + outputs += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) tl.store(c_ptrs, outputs) # for j in range(0, RANK): @@ -89,7 +85,7 @@ def split_tp_mm_kernel( # tl.store(c_ptrs, c) else: - tl.atomic_add(c_ptrs, c, sem='relaxed') + 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) @@ -101,8 +97,7 @@ def split_tp_mm_kernel( c = tl.load(c_ptrs) - buffer_ptr = tl.load(buffer_ptrs + RANK).to( - tl.pointer_type(tl.float32)) + 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) @@ -117,11 +112,9 @@ def split_tp_mm_kernel( 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.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, :]) + outputs += tl.load(buffer_ptr + offs_m[:, None] * N + offs_n[None, :]) tl.store(c_ptrs, outputs) # for j in range(0, RANK): @@ -135,10 +128,9 @@ def split_tp_mm_kernel( # tl.store(c_ptrs, c) -def triton_split_tp_gemm(x: torch.Tensor, - w: torch.Tensor, - hdl, - group: dist.ProcessGroup): +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) @@ -147,7 +139,7 @@ def triton_split_tp_gemm(x: torch.Tensor, b: right matrix with bf16 precision Returns: - c: all-reduced output + c: all-reduced output """ assert x.is_contiguous() and w.is_contiguous() M, K = x.size() @@ -165,31 +157,38 @@ def triton_split_tp_gemm(x: torch.Tensor, 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) + 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 + 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 - ) + 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/gmem_barrier_arrive_wait.py b/linghe/experimental/gmem_barrier_arrive_wait.py index b98d983..b50bad6 100644 --- a/linghe/experimental/gmem_barrier_arrive_wait.py +++ b/linghe/experimental/gmem_barrier_arrive_wait.py @@ -8,12 +8,12 @@ @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, + 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", @@ -29,13 +29,13 @@ def arrive_gmem_barrier( @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, + 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. @@ -52,8 +52,7 @@ def wait_gmem_barrier( 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. " + op == "ld" and update == 0, "Currently only support ld wait on gmem_barriers. " ) # TODO(joydddd): add support for cas barriers. diff --git a/linghe/experimental/symm_mem_barrier.py b/linghe/experimental/symm_mem_barrier.py index 666df67..04b4eb8 100644 --- a/linghe/experimental/symm_mem_barrier.py +++ b/linghe/experimental/symm_mem_barrier.py @@ -49,9 +49,9 @@ def _get_flat_tid(): @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) + tl.program_id(2) * tl.num_programs(1) * tl.num_programs(0) + + tl.program_id(1) * tl.num_programs(0) + + tl.program_id(0) ) @@ -101,12 +101,12 @@ def _wait_signal(addrs, sem: tl.constexpr): @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, + 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. @@ -156,10 +156,8 @@ def symm_mem_sync( 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") + _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 index b26c241..91f1660 100644 --- a/linghe/experimental/test_demb.py +++ b/linghe/experimental/test_demb.py @@ -2,6 +2,7 @@ """ Copyright (c) Ant Financial Service Group and its affiliates. """ + import os from datetime import timedelta @@ -10,10 +11,12 @@ 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.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 @@ -23,22 +26,24 @@ def torch_tp_emb(input_ids, weights, group): 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) + 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): +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_module.set_device(torch.device(f"cuda:{group_rank}")) - device = 'cuda' + device = "cuda" dtype = torch.bfloat16 buffers = symm_mem.empty( (B * M, D), @@ -47,26 +52,24 @@ def test_tp_emb(B=1, M=4096, N=157184, D=4096, coef=1.0, grad_coef=1.0, ) hdl = symm_mem.rendezvous(buffers, dist.group.WORLD) - local_weights = torch.randn((N, D), dtype=dtype, device=device, - requires_grad=False) + 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_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) + 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) @@ -75,50 +78,80 @@ def test_tp_emb(B=1, M=4096, N=157184, D=4096, coef=1.0, grad_coef=1.0, 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] + 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) + 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) + 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): + 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_module.set_device(torch.device(f"cuda:{group_rank}")) - device = 'cuda' + device = "cuda" dtype = torch.bfloat16 buffers = symm_mem.empty( (B * M, D), @@ -127,26 +160,24 @@ def test_sp_emb(B=1, M=4096, N=157184, D=4096, coef=1.0, grad_coef=1.0, ) hdl = symm_mem.rendezvous(buffers, dist.group.WORLD) - local_weights = torch.randn((N, D), dtype=dtype, device=device, - requires_grad=False) + 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_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) + 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)) @@ -155,54 +186,84 @@ def test_sp_emb(B=1, M=4096, N=157184, D=4096, coef=1.0, grad_coef=1.0, 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] + group_rank * N : (group_rank + 1) * N + ] global_weights.grad = None - output_ref = output_ref[group_rank * B:(group_rank + 1) * B] + 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) + 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) + 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( + 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__': + 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)) + 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) + 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) + 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 index d67f8b1..1de0151 100644 --- a/linghe/experimental/test_dla.py +++ b/linghe/experimental/test_dla.py @@ -11,8 +11,10 @@ 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.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 @@ -43,8 +45,7 @@ def torch_la(q, k, v, decay_scales, s=None, hp=False): 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) + 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) @@ -57,8 +58,7 @@ def torch_la(q, k, v, decay_scales, s=None, hp=False): 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() + 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 @@ -73,10 +73,12 @@ def rearange(x, group): 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 = 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 @@ -86,25 +88,33 @@ def select(x, group): 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] + 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): +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_module.set_device(torch.device(f"cuda:{group_rank}")) - device = torch.device('cuda') + device = torch.device("cuda") dtype = torch.bfloat16 - buffers = symm_mem.empty((B, H, 2, D, D), dtype=torch.float32, - device=device) + 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 @@ -118,8 +128,9 @@ def test_dist_la(B=1, L=4096, H=16, D=128, group=None, hpc=True, digest=False, 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 = 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_() @@ -127,8 +138,7 @@ def test_dist_la(B=1, L=4096, H=16, D=128, group=None, hpc=True, digest=False, 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, 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 @@ -141,37 +151,52 @@ def test_dist_la(B=1, L=4096, H=16, D=128, group=None, hpc=True, digest=False, 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}') + 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) + 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__': + 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)) + 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) + 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 index 07d90bf..ef31bb1 100644 --- a/linghe/experimental/test_dmm.py +++ b/linghe/experimental/test_dmm.py @@ -2,6 +2,7 @@ """ Copyright (c) Ant Financial Service Group and its affiliates. """ + import os from datetime import timedelta @@ -23,15 +24,16 @@ def torch_dist_mm(x, w, group): 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): +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_module.set_device(torch.device(f"cuda:{group_rank}")) - device = 'cuda' + device = "cuda" dtype = torch.bfloat16 buffers = symm_mem.empty((M, N), dtype=torch.float32, device=device) @@ -44,57 +46,69 @@ def test_dist_mm(M=4096, N=2048, K=4096, coef=1.0, grad_coef=1.0, # ] # buffer_tuple = tuple(buf_list) - local_weights = torch.randn((N, K), dtype=dtype, device=device, - requires_grad=False) + 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_() + 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() + 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_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) + 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__': + 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)) + 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) + 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) diff --git a/linghe/facade/emb.py b/linghe/facade/emb.py index e0bdaed..55fb5c6 100644 --- a/linghe/facade/emb.py +++ b/linghe/facade/emb.py @@ -21,19 +21,19 @@ def forward(ctx, x, w_ptr, g_ptr, dim, dtype, grad_dtype): @staticmethod def backward(ctx, grad_output): - x, = ctx.saved_tensors + (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): +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: + w_ptr: g_ptr: dim: dtype: @@ -42,9 +42,9 @@ def deprecated_fused_accumulation_embedding_lookup(x: torch.Tensor, w_ptr, lookup output """ x = x.double().requires_grad_() - return DeprecatedFusedAccumulationEmbeddingLookup.apply(x, w_ptr, g_ptr, - dim, dtype, - grad_dtype) + return DeprecatedFusedAccumulationEmbeddingLookup.apply( + x, w_ptr, g_ptr, dim, dtype, grad_dtype + ) class FusedAccumulationEmbeddingLookup(torch.autograd.Function): @@ -66,8 +66,9 @@ def backward(ctx, grad_output): return None, None, None -def fused_accumulation_embedding_lookup(x: torch.Tensor, w: torch.nn.Parameter, - grad_name: str = 'grad'): +def fused_accumulation_embedding_lookup( + x: torch.Tensor, w: torch.nn.Parameter, grad_name: str = "grad" +): """ embedding lookup Args: diff --git a/linghe/facade/fp32_gemm.py b/linghe/facade/fp32_gemm.py index 20c79cf..462bf26 100644 --- a/linghe/facade/fp32_gemm.py +++ b/linghe/facade/fp32_gemm.py @@ -5,9 +5,11 @@ import torch -from linghe.gemm.fp32_gemm import (triton_fp32_gemm, - triton_fp32_gemm_for_backward, - triton_fp32_gemm_for_update) +from linghe.gemm.fp32_gemm import ( + triton_fp32_gemm, + triton_fp32_gemm_for_backward, + triton_fp32_gemm_for_update, +) class Fp32GEMM(torch.autograd.Function): @@ -32,8 +34,7 @@ def forward(ctx, input: torch.Tensor, weight: torch.Tensor): 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]) + grad_output = grad_output.view(grad_shape[0] * grad_shape[1], grad_shape[2]) input, weight = ctx.saved_tensors diff --git a/linghe/facade/gate.py b/linghe/facade/gate.py index 1d5e02f..66a37c3 100644 --- a/linghe/facade/gate.py +++ b/linghe/facade/gate.py @@ -5,8 +5,10 @@ import torch -from linghe.utils.gate import triton_group_rms_norm_gate_forward, \ - triton_group_rms_norm_gate_backward +from linghe.utils.gate import ( + triton_group_rms_norm_gate_forward, + triton_group_rms_norm_gate_backward, +) class GroupRMSNormGateFunction(torch.autograd.Function): @@ -15,11 +17,7 @@ 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 + attn_output, gate, weight, eps=eps, group_size=group_size ) ctx.save_for_backward(attn_output, gate, weight) ctx.eps = eps @@ -32,22 +30,19 @@ 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 + dy, attn_output, gate, weight, ctx.eps, 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): +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: @@ -59,5 +54,4 @@ def group_rms_norm_gate(attn_output: torch.Tensor, Returns: output with shape [length, bs, dim] """ - return GroupRMSNormGateFunction.apply(attn_output, gate, weight, eps, - group_size) + return GroupRMSNormGateFunction.apply(attn_output, gate, weight, eps, group_size) diff --git a/linghe/facade/hadamard_quant_linear.py b/linghe/facade/hadamard_quant_linear.py index e017ea6..4bc7414 100644 --- a/linghe/facade/hadamard_quant_linear.py +++ b/linghe/facade/hadamard_quant_linear.py @@ -14,11 +14,11 @@ 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 @@ -28,18 +28,17 @@ def forward( 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) + 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 @@ -49,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) @@ -58,33 +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: @@ -99,12 +105,12 @@ class HadamardQuantLinear(torch.nn.Module): """ 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: @@ -118,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): @@ -140,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 @@ -148,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/loss.py b/linghe/facade/loss.py index d1d1abc..8f5ff29 100644 --- a/linghe/facade/loss.py +++ b/linghe/facade/loss.py @@ -5,32 +5,32 @@ import torch -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) +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, ignore_index=-100, inplace=False, - tp_group=None): + def forward(ctx, logits, labels, ignore_index=-100, inplace=False, tp_group=None): shape = logits.shape 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, labels, tp_group, - ignore_index=ignore_index) + logits, 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) + 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 @@ -48,30 +48,39 @@ def backward(ctx, grad_output): 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) + 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) + 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, None -def softmax_cross_entropy(logits: torch.Tensor, labels: torch.Tensor, - ignore_index: int = -100, inplace: bool = False, - tp_group=None): +def softmax_cross_entropy( + logits: torch.Tensor, + labels: torch.Tensor, + ignore_index: int = -100, + inplace: bool = False, + tp_group=None, +): """ softmax cross entropy Args: @@ -83,8 +92,9 @@ def softmax_cross_entropy(logits: torch.Tensor, labels: torch.Tensor, """ assert logits.is_contiguous() assert labels.is_contiguous() - return SoftmaxCrossEntropyFunction.apply(logits, labels, ignore_index, - inplace, tp_group) + return SoftmaxCrossEntropyFunction.apply( + logits, labels, ignore_index, inplace, tp_group + ) class GradScalingFunction(torch.autograd.Function): @@ -112,13 +122,15 @@ 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.save_for_backward( + logits, + ) ctx.coef = coef return loss @staticmethod def backward(ctx, grad_output): - logits, = ctx.saved_tensors + (logits,) = ctx.saved_tensors grad = triton_moe_z_loss_backward(grad_output, logits, coef=ctx.coef) return grad, None diff --git a/linghe/facade/mla.py b/linghe/facade/mla.py index 59b7640..fc440a1 100644 --- a/linghe/facade/mla.py +++ b/linghe/facade/mla.py @@ -7,27 +7,30 @@ import torch -from linghe.attn.mla import (triton_mla_forward, - triton_mla_backward, - triton_varlen_mla_forward, - triton_varlen_mla_backward) +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, - ): + 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 @@ -37,22 +40,21 @@ def forward(ctx, 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) + 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) + 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 @@ -60,43 +62,58 @@ def forward(ctx, 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) + 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, + 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, - ): +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: @@ -105,12 +122,6 @@ def multi_latend_attention(q: torch.Tensor, Returns: updated x tensor """ - return MultiLatentAttention.apply(q, - k, - v, - cu_seqlens, - padded_cu_seqlens, - max_q_length, - causal, - safe, - clip_value) + 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 be90bf7..6ec651d 100644 --- a/linghe/facade/norm.py +++ b/linghe/facade/norm.py @@ -78,8 +78,7 @@ def forward(ctx, input, weight, rms, eps, quantizer, cls, is_recomputing): 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_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, @@ -99,15 +98,13 @@ def backward(ctx, grad_output, grad_rms): 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, 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): +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 ) diff --git a/linghe/facade/permutation.py b/linghe/facade/permutation.py index d044a2e..8178c04 100644 --- a/linghe/facade/permutation.py +++ b/linghe/facade/permutation.py @@ -14,19 +14,18 @@ class _PaddedPermute(torch.autograd.Function): @staticmethod def forward( - ctx, - tokens, - probs, - routing_map, - tokens_per_expert_cuda_tensor, - tokens_per_expert_list, + ctx, + tokens, + probs, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, ): """Forward function.""" num_tokens, hidden_dim = tokens.shape row_id_map = triton_make_row_id_map(routing_map, multiple_of=16) - num_out_tokens = sum( - [(x + 15) // 16 * 16 for x in tokens_per_expert_list]) + num_out_tokens = sum([(x + 15) // 16 * 16 for x in tokens_per_expert_list]) ctx.num_tokens = num_tokens ctx.hidden_dim = hidden_dim @@ -62,11 +61,11 @@ def backward(ctx, grad_output, grad_prob, grad_map): def padded_permute( - tokens, - routing_map, - tokens_per_expert_cuda_tensor, - tokens_per_expert_list, - probs: Optional[torch.Tensor] = None, + tokens, + routing_map, + tokens_per_expert_cuda_tensor, + tokens_per_expert_list, + probs: Optional[torch.Tensor] = None, ): """Permute the tokens and probs based on the mask. Tokens with the same designated expert will be grouped together. @@ -92,8 +91,7 @@ def padded_permute( class _PaddedUnpermute(torch.autograd.Function): @staticmethod - def forward(ctx, permuted_tokens, row_id_map, tokens_per_expert, - restore_shape): + 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] @@ -107,8 +105,7 @@ def forward(ctx, permuted_tokens, row_id_map, tokens_per_expert, ctx.hidden_size = hidden_size ctx.tokens_per_expert = tokens_per_expert - output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, - None) + output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, None) return output @staticmethod @@ -129,10 +126,10 @@ def backward(ctx, grad_output): def padded_unpermute( - permuted_tokens: torch.Tensor, - row_id_map: torch.Tensor, - tokens_per_expert: torch.Tensor, - restore_shape: torch.Size, + 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 @@ -143,20 +140,19 @@ def padded_unpermute( class _BlockPaddedPermute(torch.autograd.Function): @staticmethod def forward( - ctx, - tokens, - probs, - routing_map, - tokens_per_expert_cuda_tensor, - tokens_per_expert_list, - quantizers, - cls, + 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]) + 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 ) @@ -211,13 +207,13 @@ def backward(ctx, grad_output, grad_prob, grad_map, grad_index): def block_padded_permute( - tokens, - routing_map, - tokens_per_expert_cuda_tensor, - tokens_per_expert_list, - quantizers, - cls, - probs: Optional[torch.Tensor] = None, + 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. @@ -248,15 +244,15 @@ def block_padded_permute( 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, + 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 @@ -274,8 +270,7 @@ def forward( ctx.quantizers = quantizers ctx.cls = cls - output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, - None) + output, _ = triton_unpermute_with_mask_map(permuted_tokens, row_id_map, None) return output @staticmethod @@ -309,14 +304,14 @@ def backward(ctx, grad_output): 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, + 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, diff --git a/linghe/facade/rope.py b/linghe/facade/rope.py index 0844496..b516cc6 100644 --- a/linghe/facade/rope.py +++ b/linghe/facade/rope.py @@ -7,51 +7,68 @@ import torch -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) +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, - 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): + 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) + 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 - ) + 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 @@ -70,17 +87,19 @@ def backward(ctx, grad_q, grad_k, grad_v): qkv, q_norm_weight, k_norm_weight, freqs = ctx.saved_tensors 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) + 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, @@ -98,24 +117,42 @@ def backward(ctx, grad_q, grad_k, grad_v): 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): + 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: @@ -137,39 +174,55 @@ def qk_norm_half_rope(qkv: torch.Tensor, - 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) + 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) + 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 @@ -183,40 +236,44 @@ def forward(ctx, q, kv, k_pos_emb, freqs, mscale, transpose, cu_seqlens_q, @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) + (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): +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) + 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 + 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 + 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] @@ -230,15 +287,17 @@ def mla_rope(q: torch.Tensor, - ko: shape [S, B, H, 192] or [N, H, 192] - vo: shape [S, B, H, 128] or [N, H, 128] """ - q, k, v = MLARopeFunction.apply(q, - kv, - k_pos_emb, - freqs, - mscale, - transpose, - cu_seqlens_q, - cu_seqlens_kv, - cp_size, - cp_rank, - reuse) + 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 index 02f0317..fa0f377 100644 --- a/linghe/facade/silu.py +++ b/linghe/facade/silu.py @@ -1,10 +1,11 @@ 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, - ) +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, +) class BlockSiluFunction(torch.autograd.Function): @@ -20,8 +21,8 @@ def forward(ctx, input, quantizer, grad_quantizer, cls): ctx.save_for_backward(input) x_q, x_scale, xt_q, xt_scale = triton_silu_and_block_quant_forward( - input_view, - round_scale=quantizer.force_pow_2_scales) + input_view, round_scale=quantizer.force_pow_2_scales + ) output_shape = (shape[0], shape[1], shape[2] // 2) transpose_shape = (shape[2] // 2, shape[0], shape[1]) output = cls( @@ -34,7 +35,7 @@ def forward(ctx, input, quantizer, grad_quantizer, cls): columnwise_scale_inv=xt_scale, quantizer=quantizer, requires_grad=input.requires_grad, - is_2D_scaled=False + is_2D_scaled=False, ) return output @@ -42,13 +43,12 @@ def forward(ctx, input, quantizer, grad_quantizer, cls): 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 + (input,) = ctx.saved_tensors grad_quantizer = ctx.grad_quantizer input_view = input.view(shape[0] * shape[1], shape[2] * 2) x_q, x_scale, xt_q, xt_scale = triton_silu_and_block_quant_backward( - grad_output_view, - input_view, - round_scale=grad_quantizer.force_pow_2_scales) + grad_output_view, input_view, round_scale=grad_quantizer.force_pow_2_scales + ) output = ctx.cls( shape=ctx.shape, dtype=grad_output.dtype, @@ -59,7 +59,7 @@ def backward(ctx, grad_output): columnwise_scale_inv=xt_scale, quantizer=grad_quantizer, requires_grad=ctx.input_requires_grad, - is_2D_scaled=False + is_2D_scaled=False, ) return output, None, None, None @@ -72,8 +72,17 @@ def block_silu_impl(input, quantizer, grad_quantizer, cls): class BlockBatchWeightedSiluFunction(torch.autograd.Function): @staticmethod - def forward(ctx, input, weights, counts, splits, quantizers, - grad_quantizers, cls, is_recomputing): + def forward( + ctx, + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + is_recomputing, + ): shape = input.shape ctx.grad_quantizers = grad_quantizers ctx.input_requires_grad = input.requires_grad @@ -89,17 +98,16 @@ def forward(ctx, input, weights, counts, splits, quantizers, else: output_mode = 0 - (x_q, - x_scale, - xt_q, - xt_scale) = triton_batch_weighted_silu_and_block_quant_forward(input, - weights, - counts, - splits=splits, - round_scale= - quantizers[ - 0].force_pow_2_scales, - output_mode=output_mode) + x_q, x_scale, xt_q, xt_scale = ( + triton_batch_weighted_silu_and_block_quant_forward( + input, + weights, + counts, + splits=splits, + round_scale=quantizers[0].force_pow_2_scales, + output_mode=output_mode, + ) + ) output = cls( shape=x_q.shape, @@ -111,7 +119,7 @@ def forward(ctx, input, weights, counts, splits, quantizers, columnwise_scale_inv=xt_scale, quantizer=quantizers, requires_grad=input.requires_grad, - is_2D_scaled=False + is_2D_scaled=False, ) return output @@ -119,17 +127,16 @@ def forward(ctx, input, weights, counts, splits, quantizers, def backward(ctx, grad_output): input, weights, counts = ctx.saved_tensors grad_quantizers = ctx.grad_quantizers - (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=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=grad_quantizers[0].force_pow_2_scales, + ) + ) output = ctx.cls( shape=ctx.shape, dtype=grad_output.dtype, @@ -140,21 +147,24 @@ def backward(ctx, grad_output): columnwise_scale_inv=xt_scale, quantizer=grad_quantizers, requires_grad=ctx.input_requires_grad, - is_2D_scaled=False + is_2D_scaled=False, ) return output, wgrad, None, None, None, None, None, None -def block_batch_weighted_silu_impl(input, weights, counts, splits, quantizers, - grad_quantizers, cls, is_recomputing=None): +def block_batch_weighted_silu_impl( + input, + weights, + counts, + splits, + quantizers, + grad_quantizers, + cls, + is_recomputing=None, +): assert input.ndim == 2 - output = BlockBatchWeightedSiluFunction.apply(input, - weights, - counts, - splits, - quantizers, - grad_quantizers, - cls, - is_recomputing) + output = BlockBatchWeightedSiluFunction.apply( + input, weights, counts, splits, quantizers, grad_quantizers, cls, is_recomputing + ) return output diff --git a/linghe/facade/smooth_quant_linear.py b/linghe/facade/smooth_quant_linear.py index 52cf3de..e05c756 100644 --- a/linghe/facade/smooth_quant_linear.py +++ b/linghe/facade/smooth_quant_linear.py @@ -7,8 +7,7 @@ import torch -from linghe.quant.smooth import triton_smooth_quant, \ - triton_transpose_smooth_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_transpose_and_pad @@ -16,11 +15,11 @@ 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,17 +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 @@ -53,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) @@ -61,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: @@ -109,12 +115,12 @@ class SmoothQuantLinear(torch.nn.Module): """ 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: @@ -128,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 @@ -151,10 +158,9 @@ def forward(self, input: torch.Tensor) -> torch.Tensor: 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 index 40f4809..ed56d11 100644 --- a/linghe/facade/topk.py +++ b/linghe/facade/topk.py @@ -5,10 +5,12 @@ import torch -from linghe.utils.topk import (triton_topk_forward, - triton_topk_backward, - triton_group_topk_score_forward, - triton_group_topk_score_backward) +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): @@ -24,9 +26,10 @@ def forward(ctx, x, k, dim): @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) + (indices,) = ctx.saved_tensors + grad_input = triton_topk_backward( + grad_output, indices, ctx.shape[-1], dim=ctx.dim + ) return grad_input, None, None @@ -48,15 +51,25 @@ 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) + 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 @@ -65,20 +78,21 @@ def forward(ctx, x, topk, expert_bias, num_groups, group_topk, @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) + 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'): +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: @@ -94,6 +108,6 @@ def group_topk_score(x, 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) + 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 9e6ef5f..961babf 100644 --- a/linghe/facade/transpose.py +++ b/linghe/facade/transpose.py @@ -26,7 +26,7 @@ def transpose(x, inner=True): transpose a tensor, x.ndims should not greater than 4 Args: x: input tensor - inner: + inner: if True, transpose the first two dimensions if False, transpose the last two dimensions Returns: diff --git a/linghe/gemm/blockwise_fp8_gemm.py b/linghe/gemm/blockwise_fp8_gemm.py index 6de7cd2..d2e022a 100644 --- a/linghe/gemm/blockwise_fp8_gemm.py +++ b/linghe/gemm/blockwise_fp8_gemm.py @@ -7,24 +7,23 @@ import triton import triton.language as tl - # 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, + 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) @@ -41,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 @@ -59,29 +56,40 @@ def fp8_gemm_bb_kernel( # 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): +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, +): 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 - ) + 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 @@ -93,20 +101,21 @@ def triton_bb_fp8_gemm(a: torch.Tensor, # 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, - 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, + 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, ): # a and b all tilewise quantization. pid_m = tl.program_id(axis=0) @@ -122,10 +131,8 @@ def fp8_gemm_tt_kernel( 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 = 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, :] @@ -142,24 +149,35 @@ def fp8_gemm_tt_kernel( tl.store(c_ptrs, c, mask=mask) -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): +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, - BLOCK_SIZE_M=64, - BLOCK_SIZE_N=64) + 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, + BLOCK_SIZE_M=64, + BLOCK_SIZE_N=64, + ) 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 4f213b4..ed83ef8 100644 --- a/linghe/gemm/fp32_gemm.py +++ b/linghe/gemm/fp32_gemm.py @@ -7,7 +7,6 @@ import triton import triton.language as tl - # os.environ["TRITON_PRINT_AUTOTUNING"] = "1" @@ -24,21 +23,21 @@ # @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] @@ -73,42 +72,49 @@ def triton_fp32_gemm(x: torch.Tensor, w: torch.Tensor): N, K = w.size() assert M % 32 == 0 and K % 128 == 0 and N % 16 == 0 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 + 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 = max([x for x in [16, 32, 64, 128] if N % x == 0]) num_warps = 4 num_stages = 3 - 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 - ) + 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 @@ -127,8 +133,7 @@ def fp32_gemm_for_backward_kernel( tl.store(c_ptrs, c) -def triton_fp32_gemm_for_backward(y: torch.Tensor, - w: torch.Tensor): +def triton_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): """ mix precision gemm for backward, a@b.float() Args: @@ -141,42 +146,49 @@ def triton_fp32_gemm_for_backward(y: torch.Tensor, 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 + 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, 128] if K % x == 0]) BLOCK_SIZE_M = 32 BLOCK_SIZE_N = 128 num_warps = 4 num_stages = 2 - 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 - ) + 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 @@ -208,49 +220,54 @@ def triton_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): 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 + 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 = max([x for x in [16, 32] if M % x == 0]) BLOCK_SIZE_N = 128 num_warps = 4 num_stages = 3 - 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 - ) + 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 @triton.jit def split_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, - SPLIT_COUNT: 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, + SPLIT_COUNT: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) pid_k = 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_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] + 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): @@ -265,7 +282,7 @@ def split_fp32_gemm_kernel( if SPLIT_COUNT == 1: tl.store(c_ptrs, c) else: - tl.atomic_add(c_ptrs, c, sem='relaxed') + tl.atomic_add(c_ptrs, c, sem="relaxed") def triton_split_fp32_gemm(x: torch.Tensor, w: torch.Tensor): @@ -294,50 +311,57 @@ def triton_split_fp32_gemm(x: torch.Tensor, w: torch.Tensor): 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: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"]), - SPLIT_COUNT) # noqa + 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_fp32_gemm_kernel[grid](x, w, c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - SPLIT_COUNT, - num_warps=num_warps, - num_stages=num_stages - ) + split_fp32_gemm_kernel[grid]( + x, + w, + c, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + SPLIT_COUNT, + num_warps=num_warps, + num_stages=num_stages, + ) return c # @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit def split_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, - SPLIT_COUNT: 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, + SPLIT_COUNT: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) pid_k = 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_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 + 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) @@ -353,11 +377,10 @@ def split_fp32_gemm_for_backward_kernel( if SPLIT_COUNT == 1: tl.store(c_ptrs, c) else: - tl.atomic_add(c_ptrs, c, sem='relaxed') + tl.atomic_add(c_ptrs, c, sem="relaxed") -def triton_split_fp32_gemm_for_backward(y: torch.Tensor, - w: torch.Tensor): +def triton_split_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): """ mix precision gemm for backward, a@b.float() Args: @@ -378,21 +401,28 @@ def triton_split_fp32_gemm_for_backward(y: torch.Tensor, c = torch.empty((M, N), dtype=w.dtype, device=w.device) else: c = torch.zeros((M, N), dtype=torch.float32, device=w.device) - grid = lambda META: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"]), - SPLIT_COUNT) # noqa + 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 = 2 - split_fp32_gemm_for_backward_kernel[grid](y, w, c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - SPLIT_COUNT, - num_warps=num_warps, - num_stages=num_stages - ) + split_fp32_gemm_for_backward_kernel[grid]( + y, + w, + c, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + SPLIT_COUNT, + num_warps=num_warps, + num_stages=num_stages, + ) if SPLIT_COUNT > 1: c = c.to(w.dtype) return c @@ -401,29 +431,31 @@ def triton_split_fp32_gemm_for_backward(y: torch.Tensor, # @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit def split_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, - SPLIT_COUNT: 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, + SPLIT_COUNT: tl.constexpr, ): pid_m = tl.program_id(axis=0) pid_n = tl.program_id(axis=1) pid_k = 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_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 * 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 + 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): @@ -439,7 +471,7 @@ def split_fp32_gemm_for_update_kernel( if SPLIT_COUNT == 1: tl.store(c_ptrs, c) else: - tl.atomic_add(c_ptrs, c, sem='relaxed') + tl.atomic_add(c_ptrs, c, sem="relaxed") def triton_split_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): @@ -462,19 +494,26 @@ def triton_split_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): 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: (triton.cdiv(M, META["BLOCK_SIZE_M"]), - triton.cdiv(N, META["BLOCK_SIZE_N"]), - SPLIT_COUNT) # noqa + grid = lambda META: ( + triton.cdiv(M, META["BLOCK_SIZE_M"]), + triton.cdiv(N, META["BLOCK_SIZE_N"]), + SPLIT_COUNT, + ) # noqa num_warps = 2 num_stages = 3 - split_fp32_gemm_for_update_kernel[grid](y, x, c, - M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, - SPLIT_COUNT, - num_warps=num_warps, - num_stages=num_stages - ) + split_fp32_gemm_for_update_kernel[grid]( + y, + x, + c, + M, + N, + K, + BLOCK_SIZE_K, + BLOCK_SIZE_M, + BLOCK_SIZE_N, + SPLIT_COUNT, + num_warps=num_warps, + num_stages=num_stages, + ) return c diff --git a/linghe/quant/block.py b/linghe/quant/block.py index 66cee97..1fe1905 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) @@ -28,9 +29,7 @@ def block_quant_kernel(x_ptr, y_ptr, s_ptr, M, N, BLOCK_SIZE: tl.constexpr, 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, block_size=128, round_scale=False): """ blockwise quantize x, used for blockwise recipe for weight in megatron Args: @@ -45,37 +44,45 @@ def triton_block_quant(x, assert x.is_contiguous() M, N = x.size() y = torch.empty((M, N), dtype=torch.float8_e4m3fn, device=x.device) - s = torch.empty(M // block_size, N // block_size, - dtype=torch.float32, device=x.device) + s = torch.empty( + M // block_size, 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) + 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 @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): +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, :] + 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 @@ -86,35 +93,40 @@ def blockwise_quant_kernel(x_ptr, 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) + 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) + 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) + 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): + 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: @@ -126,22 +138,19 @@ def triton_blockwise_quant(x, 2: output both Returns: - x_q: - x_scale: - xt_q: - xt_scale: + 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) + 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) + 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]( @@ -155,23 +164,24 @@ def triton_blockwise_quant(x, round_scale, output_mode, num_stages=2, - num_warps=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, - ): +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) @@ -190,41 +200,55 @@ def batch_blockwise_quant_kernel(x_ptr, 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) + 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) + 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(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) + 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): +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: @@ -234,10 +258,10 @@ def triton_batch_blockwise_quant(xs, round_scale: whether round scale to power of 2 Returns: - - x_q: - - x_scale: - - xt_q: - - xt_scale: + - x_q: + - x_scale: + - xt_q: + - xt_scale: """ assert xs.is_contiguous() @@ -248,11 +272,9 @@ def triton_batch_blockwise_quant(xs, # 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) + 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) + xt_scale = torch.empty((blocks * N,), device=device, dtype=torch.float32) if M == 0: return x_q, x_scale, xt_q, xt_scale @@ -269,7 +291,7 @@ def triton_batch_blockwise_quant(xs, n_experts, round_scale, num_stages=2, - num_warps=4 + 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 dff30e7..7394370 100644 --- a/linghe/quant/channel.py +++ b/linghe/quant/channel.py @@ -11,8 +11,9 @@ @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) @@ -21,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: @@ -54,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): @@ -85,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: @@ -99,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)) @@ -145,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 @@ -160,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] @@ -182,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) @@ -211,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 @@ -241,32 +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 cc9a4dc..effe392 100644 --- a/linghe/quant/group.py +++ b/linghe/quant/group.py @@ -10,13 +10,13 @@ @triton.jit def group_quant_kernel( - x_ptr, - y_ptr, - s_ptr, - N, - BLOCK_SIZE: tl.constexpr, - K: tl.constexpr, - ROUND: tl.constexpr, + 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) @@ -37,8 +37,7 @@ def group_quant_kernel( 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: diff --git a/linghe/quant/hadamard.py b/linghe/quant/hadamard.py index ce10a48..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): @@ -141,30 +151,12 @@ def triton_hadamard_quant(x, hm): 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 diff --git a/linghe/quant/smooth.py b/linghe/quant/smooth.py index ebac037..9c96012 100644 --- a/linghe/quant/smooth.py +++ b/linghe/quant/smooth.py @@ -12,14 +12,21 @@ @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, + max_ptr, + M, + T, + N: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, + CALIBRATE: 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, :] @@ -31,16 +38,22 @@ def tokenwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, 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) + 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 *= smooth_scale @@ -49,43 +62,57 @@ def tokenwise_smooth_quant_kernel(x_ptr, q_ptr, ss_ptr, qs_ptr, max_ptr, 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) + 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) @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, + max_ptr, + M, + N, + H: tl.constexpr, + W: tl.constexpr, + EVEN: tl.constexpr, + REVERSE: tl.constexpr, + ROUND: tl.constexpr, + CALIBRATE: 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, :] @@ -97,9 +124,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) + 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) @@ -115,8 +142,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,29 +155,31 @@ 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 + num_warps=4 if N >= 8192 else 4, ) 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 +882,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,12 +924,15 @@ 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 @@ -831,12 +953,15 @@ 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 @@ -869,13 +994,25 @@ def triton_transpose_rescale_smooth_quant(x_q, org_smooth_scale, # 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): +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) + 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 transpose: xt_q = triton_transpose_and_pad(x_q, out=xt_q, pad=pad) @@ -889,36 +1026,36 @@ def triton_smooth_quant_input(x, smooth_scale, x_q=None, x_scale=None, # 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): +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) + 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) + 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, yt_q, y_scale, yt_scale -def triton_smooth_quant_weight(w, - smooth_scale, - w_q, - quant_scale, - subrow_scales, offset=0, - round_scale=False): +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) @@ -927,17 +1064,25 @@ def triton_smooth_quant_weight(w, 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) + 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) + 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 @@ -947,21 +1092,25 @@ def triton_smooth_quant_weight(w, 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) + 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) + 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/tools/benchmark.py b/linghe/tools/benchmark.py index a112207..71ea46d 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,16 +40,22 @@ 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, + ] + ) 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)] @@ -55,17 +69,16 @@ def benchmark_func(fn, *args, n_warmup=10, n_repeat=100, ref_flops=None, 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 index 98f3ba2..460c2e3 100644 --- a/linghe/tools/check.py +++ b/linghe/tools/check.py @@ -8,8 +8,9 @@ import torch -def output_check(org_out, opt_out, name='', rtol=None, atol=None, itol=0, - amp=1.0, digest=4): +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}" @@ -20,17 +21,23 @@ def output_check(org_out, opt_out, name='', rtol=None, atol=None, itol=0, opt_dtype = opt_out.dtype if org_dtype in ( - torch.bfloat16, torch.float16, torch.float8_e4m3fn, torch.float8_e5m2): + 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): + 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): + 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): + elif org_dtype in (torch.bool, torch.uint8, torch.int8, torch.uint16, torch.int16): opt_out = opt_out.int() if rtol is None: @@ -75,10 +82,12 @@ def output_check(org_out, opt_out, name='', rtol=None, atol=None, itol=0, 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): + 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: @@ -90,16 +99,18 @@ def output_check(org_out, opt_out, name='', rtol=None, atol=None, itol=0, org_val = org_val[::itv].tolist() opt_val = opt_val[::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]) + 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]) + 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 opt: {opt_str} \n" + 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 opt: {opt_str} \n" + ) assert mismatch_count == 0, info return rel_error else: @@ -111,8 +122,10 @@ def output_check(org_out, opt_out, name='', rtol=None, atol=None, itol=0, 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" + 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 @@ -123,14 +136,16 @@ def quant_check(org_out, xq, wq, opt_out, mode): 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}') + 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=''): +def inf_or_nan(xs, name=""): if not isinstance(xs, (list, tuple)): xs = [xs] hit = False @@ -142,4 +157,5 @@ def inf_or_nan(xs, name=''): if hit: for x in xs: print( - f'{name=} {x.shape=} {x.argmax()=} {x.max()=} {x.argmin()=} {x.min()=} {x=}') + 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 c2bbf2f..33edbff 100644 --- a/linghe/tools/util.py +++ b/linghe/tools/util.py @@ -25,8 +25,9 @@ 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) @@ -54,8 +55,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 - 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))) xq = (xp / scale[:, :, None]).to(dtype) @@ -85,8 +87,7 @@ def torch_block_quant(w, B=128, dtype=torch.float8_e4m3fn, round_scale=False): w = w.clone() N, K = w.shape - wp = torch.reshape(w, (N // B, B, K // 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))) @@ -103,8 +104,7 @@ def torch_mxfp8_quant(x): 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) + 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) @@ -130,8 +130,9 @@ def torch_smooth_quant(x, smooth_scale, reverse=False, round_scale=False): 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: @@ -140,13 +141,14 @@ def torch_smooth_quant(x, smooth_scale, reverse=False, round_scale=False): return x_q, scale, x_maxs -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] @@ -166,9 +168,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) @@ -181,8 +183,12 @@ 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.float(), route_map, token_count_per_expert, indices, row_id_map @@ -246,26 +252,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 @@ -283,12 +288,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 @@ -296,24 +303,28 @@ 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 @@ -341,12 +352,14 @@ def torch_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=}') @@ -355,24 +368,28 @@ def torch_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 @@ -405,16 +422,16 @@ def fp16_f_and_b(x, w, y): 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: @@ -440,51 +457,60 @@ 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 diff --git a/linghe/utils/add.py b/linghe/utils/add.py index 8e84cb6..dbfe7c1 100644 --- a/linghe/utils/add.py +++ b/linghe/utils/add.py @@ -9,39 +9,59 @@ @triton.jit -def inplace_add_kernel(x_ptr, y_ptr, M, N, H: tl.constexpr, W: tl.constexpr, - EVEN: tl.constexpr, ACCUM: tl.constexpr): +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, :] + offs = ( + rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :] + ) if ACCUM: if EVEN: x = tl.load(x_ptr + offs) 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)) + 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), + ) else: if EVEN: y = tl.load(y_ptr + offs).to(tl.float32) 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)) + 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), + ) def triton_inplace_add(x: torch.Tensor, y: torch.Tensor, accum: bool = True): @@ -67,12 +87,6 @@ def triton_inplace_add(x: torch.Tensor, y: torch.Tensor, accum: bool = True): grid = (triton.cdiv(M, H), triton.cdiv(N, W)) inplace_add_kernel[grid]( - x, y, - M, N, - H, W, - EVEN, - accum, - num_stages=num_stages, - num_warps=num_warps + x, y, M, N, H, W, EVEN, accum, num_stages=num_stages, num_warps=num_warps ) return x diff --git a/linghe/utils/emb.py b/linghe/utils/emb.py index c91b4d5..af4edfa 100644 --- a/linghe/utils/emb.py +++ b/linghe/utils/emb.py @@ -9,19 +9,15 @@ @triton.jit -def embedding_forward_kernel(x_ptr, - y_ptr, - w_ptr, - dim, - DIM: tl.constexpr): +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) + 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): @@ -43,27 +39,14 @@ def triton_embedding_forward(x, w_ptr, dim=4096, dtype=torch.bfloat16): grid = (M,) embedding_forward_kernel[grid]( - x, - y, - w_ptr, - dim, - DIM, - num_stages=num_stages, - num_warps=num_warps + 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 + 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) @@ -76,10 +59,13 @@ def atomic_embedding_backward_kernel( else: grad_ptr = g_ptr.to(tl.pointer_type(tl.bfloat16)) - 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) + 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): @@ -115,24 +101,25 @@ def triton_atomic_embedding_backward(y, x, g_ptr, dtype=torch.bfloat16): DIM, T, num_stages=num_stages, - num_warps=num_warps + 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, - ): +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: @@ -157,12 +144,15 @@ def sync_embedding_backward_kernel(grad_output_ptr, 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) + 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) + 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): @@ -184,8 +174,7 @@ def triton_sync_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): 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) + 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 @@ -206,17 +195,14 @@ def triton_sync_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): DIM, T, num_stages=num_stages, - num_warps=num_warps + 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): +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) @@ -228,11 +214,11 @@ def scan_and_count_split_kernel(id_ptr, stop = False while not stop: min_id = tl.min(ids) - if min_id == 2 ** 30: + 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) + 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 @@ -242,28 +228,36 @@ def scan_and_count_split_kernel(id_ptr, @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): + 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) + 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) + tl.atomic_add( + accum_counts_ptr + write_index + offset + tl.arange(0, B), + counts, + mask=tl.arange(0, B) < uc, + ) write_index += uc + offset @@ -277,7 +271,14 @@ def triton_scan_and_count(ids): BLOCK = 256 assert L % BLOCK == 0 T = L // BLOCK - counts = torch.empty((B, L,), dtype=torch.int32, device=device) + 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) @@ -303,7 +304,7 @@ def triton_scan_and_count(ids): L, BLOCK, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) num_stages = 3 @@ -318,7 +319,7 @@ def triton_scan_and_count(ids): BLOCK, T, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) accum_counts = torch.cumsum(accum_counts, -1) @@ -326,10 +327,9 @@ def triton_scan_and_count(ids): @triton.jit -def deprecated_scan_and_count_kernel(id_ptr, - accum_counts_ptr, - B: tl.constexpr, - T: tl.constexpr): +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 @@ -338,7 +338,7 @@ def deprecated_scan_and_count_kernel(id_ptr, stop = False while not stop: min_id = tl.min(ids) - if min_id == 2 ** 30: + if min_id == 2**30: stop = True else: if min_id != last_min_id: @@ -346,7 +346,7 @@ def deprecated_scan_and_count_kernel(id_ptr, 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) + ids = tl.where(ids <= min_id, 2**30, ids) accum += count tl.store(accum_counts_ptr + write_index, accum) @@ -362,30 +362,26 @@ def triton_deprecated_scan_and_count(ids): T = M // B grid = (1,) deprecated_scan_and_count_kernel[grid]( - ids, - accum_counts, - B, - T, - num_stages=num_stages, - num_warps=num_warps + 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, - DIM: tl.constexpr, - T: tl.constexpr, - ): +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, + DIM: tl.constexpr, + T: tl.constexpr, +): pid = tl.program_id(axis=0).to(tl.int64) c01 = tl.load(accum_counts_ptr + pid + tl.arange(0, 2)) c0, c1 = tl.split(c01) @@ -407,12 +403,15 @@ def embedding_backward_kernel(grad_output_ptr, 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) + 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) + tl.store( + grad_ptr + input_id * dim + tl.arange(0, DIM), + outputs, + mask=tl.arange(0, DIM) < dim, + ) def triton_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): @@ -454,5 +453,5 @@ def triton_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): DIM, T, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) diff --git a/linghe/utils/gate.py b/linghe/utils/gate.py index 2829605..57977b3 100644 --- a/linghe/utils/gate.py +++ b/linghe/utils/gate.py @@ -5,42 +5,48 @@ # 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, - D: tl.constexpr, - GROUP_SIZE: tl.constexpr, - SHARE: tl.constexpr, - TRANSPOSE: tl.constexpr): +def group_rms_norm_gate_forward_kernel( + x_ptr, + gate_ptr, + weight_ptr, + out_ptr, + eps, + bs, + length, + DIM: tl.constexpr, + d: tl.constexpr, + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, + SHARE: tl.constexpr, + TRANSPOSE: 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, :] + 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) + weight_ptr + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, D), + mask=tl.arange(0, D)[None, :] < d, + ) x_offs = ( - pid * DIM + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, D)[ - None, :] + pid * 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) if TRANSPOSE: g_offs = ( - sid * bs * DIM - + bid * DIM - + tl.arange(0, GROUP_SIZE)[:, None] * d - + tl.arange(0, D)[None, :] + sid * bs * DIM + + bid * DIM + + 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) + g = tl.load(gate_ptr + g_offs, mask=tl.arange(0, D)[None, :] < d).to(tl.float32) else: g = tl.load(gate_ptr + x_offs, mask=x_offs_mask).to(tl.float32) @@ -54,12 +60,14 @@ def group_rms_norm_gate_forward_kernel(x_ptr, gate_ptr, weight_ptr, out_ptr, tl.store(out_ptr + x_offs, x, mask=x_offs_mask) -def triton_group_rms_norm_gate_forward(x: torch.Tensor, - gate: torch.Tensor, - weight: torch.Tensor, - eps=1e-6, - group_size=4, - transpose=True): +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: @@ -78,8 +86,7 @@ def triton_group_rms_norm_gate_forward(x: torch.Tensor, length, bs, dim = gate.shape else: bs, length, dim = gate.shape - assert (dim <= 8192 - and triton.next_power_of_2(group_size) == group_size) + assert dim <= 8192 and triton.next_power_of_2(group_size) == group_size assert x.is_contiguous() and gate.is_contiguous() and weight.is_contiguous() wd = weight.shape[0] share = wd != dim # all groups share the same weight @@ -116,23 +123,23 @@ def triton_group_rms_norm_gate_forward(x: torch.Tensor, @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, - DIM: tl.constexpr, - d: tl.constexpr, - D: tl.constexpr, - GROUP_SIZE: tl.constexpr, - T: tl.constexpr, - SHARE: tl.constexpr, - TRANSPOSE: tl.constexpr + grad_output_ptr, + x_ptr, + gate_ptr, + w_ptr, + dx_ptr, + dg_ptr, + dw_ptr, + eps, + bs, + length, + DIM: tl.constexpr, + d: tl.constexpr, + D: tl.constexpr, + GROUP_SIZE: tl.constexpr, + T: tl.constexpr, + SHARE: tl.constexpr, + TRANSPOSE: tl.constexpr, ): pid = tl.program_id(0) bid = pid * T // length @@ -147,17 +154,15 @@ def group_rms_norm_gate_backward_kernel( ) x_offs = ( - pid * DIM * T + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, - D)[ - None, :] + pid * DIM * T + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, D)[None, :] ) x_offs_mask = tl.arange(0, D)[None, :] < d if TRANSPOSE: offs = ( - sid * bs * DIM - + bid * DIM - + tl.arange(0, GROUP_SIZE)[:, None] * d - + tl.arange(0, D)[None, :] + 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 @@ -168,8 +173,7 @@ def group_rms_norm_gate_backward_kernel( g = tl.load(grad_output_ptr + offs, offs_mask).to(tl.float32) gate = tl.load(gate_ptr + offs, offs_mask).to(tl.float32) else: - g = tl.load(grad_output_ptr + x_offs, mask=x_offs_mask).to( - tl.float32) + g = tl.load(grad_output_ptr + x_offs, mask=x_offs_mask).to(tl.float32) gate = tl.load(gate_ptr + x_offs, mask=x_offs_mask).to(tl.float32) gate = tl.sigmoid(gate) r = tl.rsqrt(tl.sum(x * x, 1) / d + eps)[:, None] @@ -177,9 +181,8 @@ def group_rms_norm_gate_backward_kernel( dw += w_grad dx = ( - r * g * w * gate - - r * r * r * x * tl.sum(x * g * w * gate, 1, - keep_dims=True) / d + 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) @@ -196,8 +199,7 @@ def group_rms_norm_gate_backward_kernel( if SHARE: dw = tl.sum(dw, 0) - tl.store(dw_ptr + pid * d + tl.arange(0, d), dw, - mask=tl.arange(0, D) < d) + tl.store(dw_ptr + pid * d + tl.arange(0, d), dw, mask=tl.arange(0, D) < d) else: tl.store( dw_ptr @@ -209,8 +211,9 @@ def group_rms_norm_gate_backward_kernel( ) -def triton_group_rms_norm_gate_backward(grad_output, x, gate, weight, eps=1e-6, - group_size=4, transpose=True): +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: @@ -253,7 +256,7 @@ def triton_group_rms_norm_gate_backward(grad_output, x, gate, weight, eps=1e-6, share, transpose, num_stages=3, - num_warps=8 + num_warps=8, ) dw = tmp_dw.sum(dim=0) return dx, dg, dw diff --git a/linghe/utils/gather.py b/linghe/utils/gather.py index 5d7db2e..b250205 100644 --- a/linghe/utils/gather.py +++ b/linghe/utils/gather.py @@ -11,16 +11,16 @@ @triton.jit -def block_count_kernel(map_ptr, count_ptr, M, B, T: tl.constexpr, - b: tl.constexpr, E: tl.constexpr): +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 +29,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 +54,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,10 +63,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: @@ -71,10 +76,12 @@ def triton_make_row_id_map( 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 +95,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 +109,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_index_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 +142,26 @@ def make_row_id_map_and_index_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_index( - routing_map: torch.Tensor, - num_out_tokens: int, - multiple_of: int = 1, + 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 @@ -162,12 +176,15 @@ def triton_make_row_id_map_and_index( 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.zeros((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 @@ -181,7 +198,7 @@ def triton_make_row_id_map_and_index( b, n_experts, num_stages=3, - num_warps=8 + num_warps=8, ) make_row_id_map_and_index_kernel[grid]( @@ -196,14 +213,23 @@ def triton_make_row_id_map_and_index( 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 index_select_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) for i in range(T): dst_idx = pid * T + i @@ -241,32 +267,26 @@ def triton_index_select(x, indices, scale=None, out=None, scale_out=None): 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=4 + 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: @@ -274,10 +294,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): @@ -293,19 +312,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) @@ -325,13 +348,13 @@ 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 @@ -370,24 +393,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 @@ -422,34 +441,44 @@ 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_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, +): eid = tl.program_id(axis=0) cid = tl.program_id(axis=1) @@ -459,8 +488,7 @@ def batch_smooth_transpose_smooth_permute_kernel(x_ptr, scale_ptr, oss_ptr, 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: @@ -468,13 +496,14 @@ def batch_smooth_transpose_smooth_permute_kernel(x_ptr, scale_ptr, oss_ptr, 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] + 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) @@ -489,17 +518,17 @@ 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] + 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) @@ -510,16 +539,18 @@ def batch_smooth_transpose_smooth_permute_kernel(x_ptr, scale_ptr, oss_ptr, 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, + scale, + org_smooth_scale, + 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 @@ -552,8 +583,7 @@ def triton_batch_transpose_smooth_permute_with_indices(x, 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) # import pydevd @@ -577,25 +607,27 @@ def triton_batch_transpose_smooth_permute_with_indices(x, 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): +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)) @@ -607,8 +639,7 @@ def smooth_weighted_permute_with_indices_kernel(grads_ptr, 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) + 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) @@ -626,16 +657,18 @@ def smooth_weighted_permute_with_indices_kernel(grads_ptr, 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): +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: @@ -655,7 +688,7 @@ def triton_smooth_weighted_permute_with_indices(grads, assert grads.is_contiguous() M, N = grads.shape n_expert, n = smooth_scales.shape - assert N == n, f'{N=} {n=}' + assert N == n, f"{N=} {n=}" assert triton.next_power_of_2(N) == N E = indices.shape[0] device = grads.device @@ -682,25 +715,27 @@ def triton_smooth_weighted_permute_with_indices(grads, reverse, round_scale, num_stages=3, - num_warps=8 + 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 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, +): eid = tl.program_id(axis=0) wid = tl.program_id(axis=1) T = tl.num_programs(axis=1) @@ -738,15 +773,17 @@ 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_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, +): """ select and smooth and quant Args: @@ -797,24 +834,26 @@ def triton_smooth_permute_with_indices(grad_data, 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 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, +): eid = tl.program_id(axis=0) bid = tl.program_id(axis=1) n_experts = tl.num_programs(axis=0) @@ -828,11 +867,11 @@ def smooth_permute_with_mask_map_kernel(grads_data_ptr, mask = index >= 0 if index >= 0: x = tl.load(grads_data_ptr + i * N + tl.arange(0, N), mask=mask).to( - tl.float32) + tl.float32 + ) if hs > 1: - gs = tl.load(grads_scale_ptr + i * hs + tl.arange(0, hs), - mask=mask) + 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) @@ -849,21 +888,20 @@ def smooth_permute_with_mask_map_kernel(grads_data_ptr, 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) + 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 + 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, ): """ gather ( and optional dequant) and smooth quant @@ -886,9 +924,11 @@ def triton_smooth_permute_with_mask_map( assert inp.is_contiguous() assert row_id_map.shape[1] == num_experts assert triton.next_power_of_2(hidden_size) == hidden_size - output = torch.empty((num_out_tokens, hidden_size), - dtype=torch.float8_e4m3fn, - device=row_id_map.device) + output = torch.empty( + (num_out_tokens, hidden_size), + dtype=torch.float8_e4m3fn, + device=row_id_map.device, + ) if scale is None: hs = 0 else: @@ -912,26 +952,27 @@ def triton_smooth_permute_with_mask_map( hidden_size, hs, reverse, - round_scale + round_scale, ) return output, permuted_scale @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 - ): +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) @@ -953,9 +994,10 @@ def batch_block_pad_permute_with_indices_kernel(x_ptr, 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) + 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: @@ -965,39 +1007,48 @@ def batch_block_pad_permute_with_indices_kernel(x_ptr, tl.store( xs_ptr + psi * nb + cid * padding_count + rid * 128 + tl.arange(0, 128), 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) + 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(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) + 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) + 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): +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 Args: @@ -1009,11 +1060,11 @@ def triton_batch_block_pad_permute_with_indices(xs, round_scale: whether round scale to power of 2 Returns: - x_q: - x_scale: - xt_q: - xt_scale: - prob_output: + x_q: + x_scale: + xt_q: + xt_scale: + prob_output: """ assert xs.is_contiguous() @@ -1056,8 +1107,7 @@ def triton_batch_block_pad_permute_with_indices(xs, round_scale, PROB, num_stages=2, - num_warps=4 + num_warps=4, ) return x_q, x_scale, xt_q, xt_scale, prob_output - diff --git a/linghe/utils/loss.py b/linghe/utils/loss.py index 64c3ab6..cf63b15 100644 --- a/linghe/utils/loss.py +++ b/linghe/utils/loss.py @@ -9,14 +9,16 @@ @triton.jit -def softmax_cross_entropy_forward_kernel(logit_ptr, - label_ptr, - loss_ptr, - sum_exp_ptr, - max_logit_ptr, - N, - ignore_index, - 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: @@ -30,13 +32,16 @@ def softmax_cross_entropy_forward_kernel(logit_ptr, 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) + 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)) + tl.exp(logit - latest_max_logit) + ) max_logit = latest_max_logit tl.store(sum_exp_ptr + pid, sum_exp) @@ -74,20 +79,24 @@ def triton_softmax_cross_entropy_forward(logits, labels, ignore_index=-100): ignore_index, B, num_stages=3, - num_warps=2 + 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, - output_grad_ptr, - input_grad_ptr, - N, - ignore_index, - B: tl.constexpr, - INPLACE: 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) @@ -95,12 +104,17 @@ def softmax_cross_entropy_backward_kernel(logit_ptr, label_ptr, sum_exp_ptr, 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) + 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.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) @@ -111,16 +125,24 @@ def softmax_cross_entropy_backward_kernel(logit_ptr, label_ptr, sum_exp_ptr, 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=-1e10).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 if INPLACE: - tl.store(logit_ptr + pid * N + i * B + tl.arange(0, B), grad, - mask=i * B + tl.arange(0, B) < N) + 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.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) @@ -128,10 +150,9 @@ def softmax_cross_entropy_backward_kernel(logit_ptr, label_ptr, sum_exp_ptr, tl.store(input_grad_ptr + pid * N + label, target_grad) -def triton_softmax_cross_entropy_backward(logits, labels, sum_exp, max_logit, - output_grad, - ignore_index=-100, - inplace=False): +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: @@ -166,7 +187,7 @@ def triton_softmax_cross_entropy_backward(logits, labels, sum_exp, max_logit, B, inplace, num_stages=3, - num_warps=8 + num_warps=8, ) if inplace: dx = logits @@ -174,16 +195,18 @@ def triton_softmax_cross_entropy_backward(logits, labels, sum_exp, max_logit, @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): +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) @@ -198,13 +221,16 @@ def parallel_logit_stat_kernel(logit_ptr, 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) + 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)) + tl.exp(logit - latest_max_logit) + ) max_logit = latest_max_logit tl.store(sum_exp_ptr + pid, sum_exp) @@ -213,17 +239,22 @@ def parallel_logit_stat_kernel(logit_ptr, if label // N == group_rank: target_logit = tl.load(logit_ptr + pid * N + label % N).to(tl.float32) else: - target_logit = float('-inf') + 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): +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: @@ -235,14 +266,15 @@ def parallel_calc_loss_kernel(label_ptr, stats, sum_exp_ptr, max_logit_ptr, sum_exp = 0.0 sum_exp = sum_exp.to(tl.float64) max_logit = -1e9 - tg = float('-inf') # target logit + 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) + ml - latest_max_logit + ) max_logit = latest_max_logit loss = tl.log(sum_exp) - (tg - max_logit) @@ -257,8 +289,9 @@ def parallel_calc_loss_kernel(label_ptr, stats, sum_exp_ptr, max_logit_ptr, """ -def triton_parallel_softmax_cross_entropy_forward(logits, labels, group, - ignore_index=-100): +def triton_parallel_softmax_cross_entropy_forward( + logits, labels, group, ignore_index=-100 +): """ compute token-wise softmax cross entropy loss Args: @@ -279,8 +312,7 @@ def triton_parallel_softmax_cross_entropy_forward(logits, labels, group, sum_exp = stats[0] max_logit = stats[1] target_logit = stats[2] - statistic = torch.empty((3 * group_size, M), device=device, - dtype=torch.float32) + statistic = torch.empty((3 * group_size, M), device=device, dtype=torch.float32) B = 2048 grid = (M,) parallel_logit_stat_kernel[grid]( @@ -295,32 +327,41 @@ def triton_parallel_softmax_cross_entropy_forward(logits, labels, group, group_size, B, num_stages=3, - num_warps=2 + 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) + 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): +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) @@ -329,12 +370,17 @@ def parallel_softmax_cross_entropy_backward_kernel(logit_ptr, label_ptr, 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) + 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.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) @@ -343,34 +389,45 @@ def parallel_softmax_cross_entropy_backward_kernel(logit_ptr, label_ptr, 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) + 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) + 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.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 + 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): +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: @@ -411,7 +468,7 @@ def triton_parallel_softmax_cross_entropy_backward(logits, labels, sum_exp, B, inplace, num_stages=3, - num_warps=8 + num_warps=8, ) if inplace: dx = logits @@ -419,15 +476,14 @@ def triton_parallel_softmax_cross_entropy_backward(logits, labels, sum_exp, @triton.jit -def moe_z_loss_forward_kernel(logit_ptr, loss_ptr, coef, - T: tl.constexpr, - D: tl.constexpr): +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) + 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) @@ -457,31 +513,24 @@ def triton_moe_z_loss_forward(logits, coef=1e-6): 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 - ) + 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): +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) + 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) @@ -489,8 +538,10 @@ def moe_z_loss_backward_kernel(input_grad_ptr, logit_ptr, output_grad_ptr, coef, 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) + 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): @@ -518,13 +569,6 @@ def triton_moe_z_loss_backward(grads, logits, coef=1e-6): 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 + 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 index fc4fb67..f324001 100644 --- a/linghe/utils/mul.py +++ b/linghe/utils/mul.py @@ -47,13 +47,7 @@ def triton_dot(x, y): 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 - ) + dot_kernel[grid](x, y, s, M, N, H, W, num_stages=num_stages, num_warps=num_warps) return s @@ -80,22 +74,19 @@ def triton_inplace_scale(x, scale): B = 512 m = x.numel() grid = (triton.cdiv(m, B),) - inplace_scale_kernel[grid]( - x, - scale, - m, - B, - num_stages=2, - num_warps=2 - ) + 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, ): +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) @@ -110,7 +101,7 @@ def batch_scale_kernel(input_ptrs, size_ptr, scale, 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) + 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) @@ -134,10 +125,12 @@ def triton_batch_scale(xs, scale): 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) + 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 @@ -145,14 +138,5 @@ def triton_batch_scale(xs, scale): 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 - ) + 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 d2c2689..77133a8 100644 --- a/linghe/utils/norm.py +++ b/linghe/utils/norm.py @@ -11,42 +11,47 @@ @triton.jit -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): +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), - mask=tl.arange(0, N) < 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): 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) + 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) + 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) + 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 - tl.store(out_ptr + offs, x, - mask=mask) + tl.store(out_ptr + offs, x, mask=mask) offs += n * W @@ -57,7 +62,7 @@ def triton_rms_norm_forward(x, weight, eps=1e-6, out=None, rms=None): 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, + rms: use x*rms to calculate output if rms is not None, it will accelerate recompute of rms norm Returns: out: output tensor @@ -83,70 +88,57 @@ def triton_rms_norm_forward(x, weight, eps=1e-6, out=None, rms=None): grid = (triton.cdiv(M, T * W),) rms_norm_forward_kernel[grid]( - x, - weight, - out, - rms, - eps, - M, - T, - n, - N, - W, - REUSE, - 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, rms @triton.jit def rms_norm_backward_kernel( - 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 + 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), mask=tl.arange(0, N) < 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 * W * T + i * W + tl.arange(0, W)[:, None] < M) & ( - tl.arange(0, N) < n) + 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) 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] + 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 - tl.store(dw_ptr + pid * n + tl.arange(0, N), w_grads, - mask=tl.arange(0, N) < n) + 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, rms=None): @@ -182,7 +174,7 @@ def triton_rms_norm_backward(grad_output, x, w, eps=1e-6, rms=None): W, REUSE, num_stages=3, - num_warps=4 + num_warps=4, ) return dx, tmp_dw.sum(dim=0) @@ -190,31 +182,31 @@ def triton_rms_norm_backward(grad_output, x, w, eps=1e-6, rms=None): # output non-transposed and transposed together # 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, - n, - N: tl.constexpr, - T: 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) 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, :] + 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) masks = (indices[:, None] < M) & (tl.arange(0, N) < n) @@ -228,16 +220,16 @@ def rms_norm_and_block_quant_forward_kernel(x_ptr, scale = tl.exp2(tl.ceil(tl.log2(scale))) x = (x / scale[:, :, None]).to(out_ptr.dtype.element_ty) x = tl.reshape(x, [W, N]) - 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( + 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] @@ -251,35 +243,34 @@ def rms_norm_and_block_quant_forward_kernel(x_ptr, 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) + 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, - n, - N: tl.constexpr, - T: 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) 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, :] + 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) masks = (indices[:, None] < M) & (tl.arange(0, N) < n) @@ -297,34 +288,43 @@ def rms_norm_and_block_quant_forward_n_kernel(x_ptr, x = x / scale[:, :, None] x = tl.reshape(x, [W, N]) - 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( + 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) @@ -340,15 +340,16 @@ def rms_norm_and_block_quant_forward_t_kernel(x_ptr, 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: @@ -385,10 +386,10 @@ def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, 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) + 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) @@ -409,7 +410,7 @@ def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, W, round_scale, num_stages=3, - num_warps=4 + num_warps=4, ) elif output_mode == 1: # only output transposed tensor @@ -417,17 +418,19 @@ def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, W = 32 assert n % W == 0 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) + 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 # we force set output_mode=2 when recompute qkv, but it has rms @@ -456,7 +459,7 @@ def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, H, round_scale, num_stages=2, - num_warps=4 + num_warps=4, ) else: W = 8192 // N @@ -476,40 +479,47 @@ def triton_rms_norm_and_block_quant_forward(x: torch.Tensor, W, round_scale, num_stages=3, - num_warps=4 + num_warps=4, ) W = 32 - assert n % W == 0, f' {n=} {W=}' + 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) + 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, + ) return out, scale, rms, transpose_output, transpose_scale @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): +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, :] @@ -518,8 +528,7 @@ def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, 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) @@ -545,12 +554,18 @@ def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, # 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, + calibrate=False, + output_rms=False, + round_scale=False, +): """""" assert x.is_contiguous() and weight.is_contiguous() M, N = x.shape @@ -590,31 +605,32 @@ def triton_rms_norm_and_smooth_quant_forward(x, weight, smooth_scale=None, 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 + @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 + 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) @@ -639,7 +655,8 @@ def rms_norm_fp32_gemm_block_quant_forward_n_kernel( 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) + tl.float32 + ) x = tl.load(x_ptr + x_offs).to(tl.float32) w = tl.load(w_ptrs).to(tl.float32) @@ -655,8 +672,8 @@ def rms_norm_fp32_gemm_block_quant_forward_n_kernel( x = x / scale[:, None] tl.store( - xs_ptr + M * i + pid_m * BLOCK_SIZE_M + tl.arange(0, BLOCK_SIZE_M), - scale) + 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 @@ -668,15 +685,15 @@ def rms_norm_fp32_gemm_block_quant_forward_n_kernel( 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 - ): +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 @@ -694,12 +711,16 @@ def triton_rms_norm_fp32_gemm_block_quant_forward(x: torch.Tensor, - y: rms normed tensor - rms: 1/rms - logits: router logit - - x_q: - - x_s: + - x_q: + - x_s: - xt_q: - xt_s: """ - assert x.is_contiguous() and norm_weight.is_contiguous() and route_weight.is_contiguous() + 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() @@ -723,36 +744,41 @@ def triton_rms_norm_fp32_gemm_block_quant_forward(x: torch.Tensor, 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 - ) + 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) + 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 89af89f..ba178c9 100644 --- a/linghe/utils/rearange.py +++ b/linghe/utils/rearange.py @@ -9,11 +9,20 @@ @triton.jit -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): +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) @@ -28,10 +37,15 @@ def sort_chunks_by_index_kernel(x_ptr, y_ptr, scale_ptr, scale_output_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_sort_chunks_by_index(x, counts, indices, scales=None): @@ -78,6 +92,6 @@ def triton_sort_chunks_by_index(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 24eb767..40960c8 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 @@ -77,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 @@ -117,33 +122,25 @@ def triton_batch_count_zero(xs): """ 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).cuda(device, non_blocking=True) - ptrs = torch.tensor([x.data_ptr() for x in xs], - dtype=torch.int64).cuda(device, non_blocking=True) + 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 + ) block = 2048 tensor_count = len(xs) - counts = torch.empty((tensor_count, block), device=device, - dtype=torch.int64) + 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 - ) + batch_count_zero_kernel[grid](ptrs, sizes, counts, B, num_stages=2, num_warps=2) count = counts.sum() return count @triton.jit -def norm_kernel(input_ptr, tmp_ptr, m, - 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) @@ -166,7 +163,7 @@ def triton_norm(x, ord=2, norm=True, scalar=True): 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) + 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: @@ -181,15 +178,7 @@ def triton_norm(x, ord=2, norm=True, scalar=True): T = triton.cdiv(m, B) tmp = torch.empty((T,), device=device, dtype=torch.float32) grid = (T,) - norm_kernel[grid]( - x, - tmp, - m, - B, - ord, - num_stages=2, - num_warps=2 - ) + norm_kernel[grid](x, tmp, m, B, ord, num_stages=2, num_warps=2) if ord == -1: output = tmp.max() else: @@ -202,11 +191,15 @@ def triton_norm(x, ord=2, norm=True, scalar=True): @triton.jit -def batch_norm_kernel(input_ptrs, size_ptr, tmp_ptr, - DT: tl.constexpr, - B: tl.constexpr, - ORD: tl.constexpr, - HP: tl.constexpr): +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).to(tl.int64) sm = tl.num_programs(axis=1) @@ -251,43 +244,39 @@ def triton_batch_norm(xs, ord=2, norm=True, scalar=True, high_precision=True): 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) + 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 """ if len(xs) == 0: - return torch.zeros(() if scalar else (1,), device='cuda', - dtype=torch.float32) + 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).cuda(device, non_blocking=True) - ptrs = torch.tensor([x.data_ptr() for x in xs], - dtype=torch.int64).cuda(device, non_blocking=True) + 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 sm = 256 tensor_count = len(xs) - tmp = torch.empty((tensor_count, sm), device=device, - dtype=torch.float64 if high_precision else torch.float32) + tmp = torch.empty( + (tensor_count, sm), + device=device, + dtype=torch.float64 if high_precision else torch.float32, + ) B = 128 grid = (tensor_count, sm) batch_norm_kernel[grid]( - ptrs, - sizes, - tmp, - DT, - B, - ord, - high_precision, - num_stages=2, - num_warps=2 + ptrs, sizes, tmp, DT, B, ord, high_precision, num_stages=2, num_warps=2 ) if ord == -1: output = tmp.max() diff --git a/linghe/utils/rope.py b/linghe/utils/rope.py index e89a5ae..5f70afc 100644 --- a/linghe/utils/rope.py +++ b/linghe/utils/rope.py @@ -9,15 +9,21 @@ @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) @@ -30,119 +36,163 @@ def half_rope_forward_kernel(q_ptr, k_ptr, freqs_ptr, qo_ptr, ko_ptr, 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): @@ -179,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, @@ -191,20 +243,23 @@ 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) @@ -217,82 +272,93 @@ 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, :]).to(tl.float32) + 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, :]).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)) + 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, :], qo) + 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, :], qo) + 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, :]).to(tl.float32) + 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, :]).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)) + 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, :], ko) + 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, :], ko) + 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): +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: @@ -306,7 +372,8 @@ def triton_half_rope_backward(q_grad, k_grad, freqs, inplace=False, grid = (L,) half_rope_backward_kernel[grid]( - q_grad, k_grad, + q_grad, + k_grad, freqs, B, H, @@ -315,26 +382,31 @@ def triton_half_rope_backward(q_grad, k_grad, freqs, inplace=False, 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, - SILU: 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, + 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 @@ -358,22 +430,36 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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, :]).to(tl.float32) + q_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) q1 = tl.load( - q_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: q0 = tl.load( - q_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) q1 = tl.load( - q_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: q0 = q0 * tl.sigmoid(q0) q1 = q1 * tl.sigmoid(q1) @@ -381,21 +467,34 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, q1 *= rms[:, None] q1 *= q_weight_1 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)[:, None] + + tl.arange(0, D)[None, :], + q1, + ) 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, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (H, 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) + 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)).to(tl.float32) k_weight_1 = tl.load(k_norm_weight_ptr + D + tl.arange(0, D)).to(tl.float32) @@ -408,22 +507,36 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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, :]).to(tl.float32) + k_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: k0 = tl.load( - k_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) k1 = tl.load( - k_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: k0 = k0 * tl.sigmoid(k0) k1 = k1 * tl.sigmoid(k1) @@ -431,21 +544,34 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, k1 *= rms[:, None] k1 *= k_weight_1 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)[:, None] + + tl.arange(0, D)[None, :], + k1, + ) 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, 2, d)), (0, 2, 1)), dim=2) + * signs, + (0, 2, 1), + ), + (h, 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)[:, None] + + tl.arange(0, D)[None, :], + k0, + ) if INTERLEAVED: row_offs = tl.arange(0, h) * (w + 2) @@ -456,55 +582,81 @@ def qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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, :]).to(tl.float32) + v_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) v1 = tl.load( - v_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: v0 = tl.load( - v_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) v1 = tl.load( - v_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: v0 = v0 * tl.sigmoid(v0) v1 = v1 * tl.sigmoid(v1) tl.store( - vo_ptr + pid * h * DD + i * L * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :], v0) + vo_ptr + + pid * h * DD + + i * L * h * DD + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + v0, + ) tl.store( - vo_ptr + pid * h * DD + i * L * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :], v1) + vo_ptr + + pid * h * DD + + i * L * h * DD + + D + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + v1, + ) @triton.jit -def compatible_qk_norm_and_half_rop_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): +def compatible_qk_norm_and_half_rop_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 @@ -532,22 +684,40 @@ def compatible_qk_norm_and_half_rop_forward_kernel(qkv_ptr, 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, :], mask=row_mask).to(tl.float32) + 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) + 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, :], mask=row_mask).to(tl.float32) + 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) + 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)) @@ -556,23 +726,36 @@ def compatible_qk_norm_and_half_rop_forward_kernel(qkv_ptr, 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_p)[ - :, - None] + tl.arange( - 0, D)[None, :], q1, mask=q_mask) + 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_p, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (H_p, 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_p)[:, - None] + tl.arange(0, - D)[ - None, :], q0, - mask=q_mask) + 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) @@ -590,22 +773,40 @@ def compatible_qk_norm_and_half_rop_forward_kernel(qkv_ptr, 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, :], mask=row_mask).to(tl.float32) + 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, :], mask=row_mask).to(tl.float32) + 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, :], mask=row_mask).to(tl.float32) + 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, :], mask=row_mask).to(tl.float32) + 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) @@ -615,25 +816,35 @@ def compatible_qk_norm_and_half_rop_forward_kernel(qkv_ptr, 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_p)[ - :, - None] + tl.arange( - 0, D)[None, :], k1, - mask=k_mask + 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_p, 2, d)), (0, 2, 1)), - dim=2) * signs, (0, 2, 1)), (h_p, 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_p)[:, - None] + tl.arange(0, - D)[ - None, :], k0, - mask=k_mask + 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: @@ -650,45 +861,78 @@ def compatible_qk_norm_and_half_rop_forward_kernel(qkv_ptr, 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, :], mask=row_mask).to(tl.float32) + 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) + 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) + 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) + 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 + DD * tl.arange(0, h_p)[:, - None] + tl.arange(0, - D)[ - None, :], v0, - mask=v_mask) + 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) + 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 - ): +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 @@ -725,8 +969,7 @@ def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, 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=}' + 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) @@ -743,9 +986,12 @@ def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, if H_p == H and h_p == 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, @@ -757,14 +1003,17 @@ def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, transposed, silu, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) else: compatible_qk_norm_and_half_rop_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, @@ -778,32 +1027,35 @@ def triton_qk_norm_and_half_rope_forward(qkv, q_norm_weight, k_norm_weight, 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, - SILU: 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, + 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 @@ -829,39 +1081,63 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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, :]).to( - tl.float32) + gq_ptr + + i * L * H * DD + + pid * H * DD + + DD * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ).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, :]).to(tl.float32) + gq_ptr + + i * L * H * DD + + pid * H * DD + + D + + DD * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) - 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_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_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, :]).to(tl.float32) + q_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) q1 = tl.load( - q_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: q0 = tl.load( - q_ptr + pid * stride + i * L * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q_ptr + + pid * stride + + i * L * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) q1 = tl.load( - q_ptr + pid * stride + i * L * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q_ptr + + pid * stride + + i * L * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: s0 = tl.sigmoid(q0) @@ -869,9 +1145,9 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -885,8 +1161,7 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -898,26 +1173,40 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, if TRANSPOSED: tl.store( - dq_ptr + pid * B * grad_stride + i * grad_stride + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dq_0) + 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) + dq_ptr + + pid * B * grad_stride + + i * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_1, + ) else: tl.store( - dq_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dq_0) + 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) + dq_ptr + + pid * grad_stride + + i * L * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_1, + ) 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) @@ -938,39 +1227,63 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, # [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, :]).to( - tl.float32) + gk_ptr + + i * L * h * DD + + pid * h * DD + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).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, :]).to(tl.float32) + gk_ptr + + i * L * h * DD + + pid * h * DD + + D + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) - 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_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_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, :]).to(tl.float32) + k_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: k0 = tl.load( - k_ptr + pid * stride + i * L * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + + pid * stride + + i * L * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * stride + i * L * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + + pid * stride + + i * L * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: @@ -979,9 +1292,9 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -995,8 +1308,7 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -1008,26 +1320,40 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, if TRANSPOSED: tl.store( - dk_ptr + pid * B * grad_stride + i * grad_stride + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dk_0) + 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) + dk_ptr + + pid * B * grad_stride + + i * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_1, + ) else: tl.store( - dk_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dk_0) + 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) + dk_ptr + + pid * grad_stride + + i * L * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_1, + ) 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) @@ -1043,35 +1369,54 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, for i in range(B): gv_0 = tl.load( - gv_ptr + i * L * h * DD + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :]).to( - tl.float32) + gv_ptr + + i * L * h * DD + + pid * h * DD + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) gv_1 = tl.load( - gv_ptr + i * L * h * DD + pid * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + gv_ptr + + i * L * h * DD + + pid * h * DD + + D + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).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, :]).to(tl.float32) + v_ptr + + pid * B * stride + + i * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) v1 = tl.load( - v_ptr + pid * B * stride + i * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + + pid * B * stride + + i * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) else: v0 = tl.load( - v_ptr + i * L * stride + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + + i * L * stride + + pid * stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) v1 = tl.load( - v_ptr + i * L * stride + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + + i * L * stride + + pid * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) s0 = tl.sigmoid(v0) s1 = tl.sigmoid(v1) @@ -1083,50 +1428,68 @@ def qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, if TRANSPOSED: tl.store( - dv_ptr + pid * B * grad_stride + i * grad_stride + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dv_0) + dv_ptr + + pid * B * grad_stride + + i * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_0, + ) tl.store( - dv_ptr + pid * B * grad_stride + i * grad_stride + D + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dv_1) + dv_ptr + + pid * B * grad_stride + + i * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_1, + ) else: tl.store( - dv_ptr + pid * grad_stride + i * L * grad_stride + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dv_0) + dv_ptr + + pid * grad_stride + + i * L * grad_stride + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_0, + ) tl.store( - dv_ptr + pid * grad_stride + i * L * grad_stride + D + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dv_1) + dv_ptr + + pid * grad_stride + + i * L * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_1, + ) @triton.jit -def compatible_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 - ): +def compatible_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 @@ -1156,48 +1519,74 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, for i in range(B): gq_0 = tl.load( - 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 + 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_p)[ - :, - None] + tl.arange( - 0, D)[None, :] - , mask=tl.arange(0, H_p)[:, None] < H + 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_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, :]) # 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) + 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) + 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, :]) # q1 = tl.load(q_ptr + pid * stride + i * L * stride + D + 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) + 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, :], mask=row_mask).to(tl.float32) + q_ptr + + pid * stride + + i * L * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) if SILU: s0 = tl.sigmoid(q0) @@ -1205,9 +1594,9 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -1221,8 +1610,7 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -1236,29 +1624,47 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, # 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) + 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) + 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, mask=row_mask) + 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) + 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) @@ -1283,43 +1689,69 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, # [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_p)[:, - None] + tl.arange(0, - D)[ - None, :], - mask=tl.arange(0, h_p)[:, None] < h + 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_p)[ - :, - None] + tl.arange( - 0, D)[None, :], - mask=tl.arange(0, h_p)[:, None] < h + 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_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, :], mask=row_mask).to(tl.float32) + 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, :], mask=row_mask).to(tl.float32) + 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, :], mask=row_mask).to(tl.float32) + 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, :], mask=row_mask).to(tl.float32) + k_ptr + + pid * stride + + i * L * stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) if SILU: @@ -1328,9 +1760,9 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -1344,8 +1776,7 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -1357,26 +1788,44 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, if TRANSPOSED: 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) + 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) + 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, mask=row_mask) + 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) + 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) @@ -1397,37 +1846,60 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, for i in range(B): 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_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) + 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) + 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) + 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) + 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) + 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) @@ -1439,32 +1911,59 @@ def compatible_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, if TRANSPOSED: 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) + 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) + 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 + DD * row_offs[ - :, - None] + tl.arange( - 0, D)[None, :], dv_0, mask=row_mask) + 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): + 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: @@ -1535,17 +2034,21 @@ def triton_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, transposed, silu, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) else: compatible_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, @@ -1560,7 +2063,7 @@ def triton_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, transposed, silu, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) dqw = tmp_dqw.sum(0) dkw = tmp_dkw.sum(0) @@ -1568,12 +2071,16 @@ def triton_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, @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 +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) + cun = tl.min(tl.where(cus <= cu, 2**24, cus), 0) length = cun - cu token_idx = pid_m - cu @@ -1582,15 +2089,14 @@ def _get_varlen_token_idx(cu_seqlens, pid_m, seq_num, padded_seq_num, cp_rank, token_idx = token_idx + cp_rank * length // 2 else: token_idx = (token_idx - length // 2) + ( - 2 * cp_size - cp_rank - 1 + 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): +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 @@ -1600,35 +2106,36 @@ def _get_fixlen_token_idx(num_tokens, pid_m, seq_num, cp_rank, cp_size, 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 + 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, - D: tl.constexpr, - d: tl.constexpr, - INTERLEAVED: tl.constexpr, - SILU: tl.constexpr, - CP_SIZE: tl.constexpr, - REUSE: tl.constexpr - ): +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, + 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) @@ -1651,12 +2158,12 @@ def varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, else: row_offs = tl.arange(0, H) - q0 = tl.load(q_ptr + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) - q1 = tl.load(q_ptr + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q0 = tl.load( + q_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) + q1 = tl.load( + q_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: q0 = q0 * tl.sigmoid(q0) @@ -1665,28 +2172,37 @@ def varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, q1 *= rms[:, None] q1 *= q_weight_1 tl.store( - qo_ptr + pid * H * DD + D + DD * tl.arange(0, H)[:, - None] + tl.arange( - 0, D)[None, :], q1) + qo_ptr + + pid * H * DD + + D + + DD * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :], + q1, + ) 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, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (H, D), + ) q0 = q0 * cos + qr * sin tl.store( - qo_ptr + pid * H * DD + DD * tl.arange(0, H)[:, - None] + tl.arange(0, - D)[ - None, :], q0) + qo_ptr + + pid * 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)).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) + 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 @@ -1698,13 +2214,12 @@ def varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, row_offs = tl.arange(0, h) k_ptr = qkv_ptr + DD * H - k0 = tl.load(k_ptr + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k0 = tl.load( + k_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: k0 = k0 * tl.sigmoid(k0) @@ -1713,21 +2228,31 @@ def varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, k1 *= rms[:, None] k1 *= k_weight_1 tl.store( - ko_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :], k1) + ko_ptr + + pid * h * DD + + D + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + k1, + ) 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, 2, d)), (0, 2, 1)), dim=2) * signs, + (0, 2, 1), + ), + (h, D), + ) k0 = k0 * cos + kr * sin tl.store( - ko_ptr + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :], k0) + ko_ptr + + pid * h * DD + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + k0, + ) if INTERLEAVED: row_offs = tl.arange(0, h) * (w + 2) @@ -1736,55 +2261,62 @@ def varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, row_offs = tl.arange(0, 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, :]).to(tl.float32) + v0 = tl.load( + v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) v1 = tl.load( - v_ptr + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: v0 = v0 * tl.sigmoid(v0) v1 = v1 * tl.sigmoid(v1) tl.store( - vo_ptr + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :], v0) + vo_ptr + + pid * h * DD + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + v0, + ) tl.store( - vo_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :], v1) + vo_ptr + + pid * h * DD + + D + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :], + v1, + ) @triton.jit -def compatible_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 - ): +def compatible_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) @@ -1810,12 +2342,14 @@ def compatible_varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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) + 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) @@ -1824,29 +2358,39 @@ def compatible_varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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) + 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)) + 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) + 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) + 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 @@ -1860,13 +2404,14 @@ def compatible_varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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) + 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) + 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) @@ -1876,22 +2421,33 @@ def compatible_varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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) + 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)) + 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) + 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) @@ -1902,13 +2458,14 @@ def compatible_varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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) + 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) + 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) @@ -1916,28 +2473,41 @@ def compatible_varlen_qk_norm_and_half_rope_forward_kernel(qkv_ptr, 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) + 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 - ): + 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 @@ -1981,11 +2551,14 @@ def triton_varlen_qk_norm_and_half_rope_forward(qkv, q_norm_weight, if PH == H and ph == h: varlen_qk_norm_and_half_rope_forward_kernel[grid]( qkv, - q_norm_weight, k_norm_weight, + q_norm_weight, + k_norm_weight, freqs, cu_seqlens_q, cu_seqlens_kv, - qo, ko, vo, + qo, + ko, + vo, stride, eps, mscale, @@ -2001,16 +2574,19 @@ def triton_varlen_qk_norm_and_half_rope_forward(qkv, q_norm_weight, cp_size, reuse, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) else: compatible_varlen_qk_norm_and_half_rope_forward_kernel[grid]( qkv, - q_norm_weight, k_norm_weight, + q_norm_weight, + k_norm_weight, freqs, cu_seqlens_q, cu_seqlens_kv, - qo, ko, vo, + qo, + ko, + vo, stride, eps, mscale, @@ -2028,37 +2604,41 @@ def triton_varlen_qk_norm_and_half_rope_forward(qkv, q_norm_weight, cp_size, reuse, num_stages=num_stages, - num_warps=num_warps + 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, - D: tl.constexpr, - d: tl.constexpr, - INTERLEAVED: tl.constexpr, - SILU: tl.constexpr, - CP_SIZE: tl.constexpr, - REUSE: tl.constexpr - ): +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, + 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 @@ -2084,27 +2664,31 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, row_offs = tl.arange(0, H) gq_0 = tl.load( - gq_ptr + pid * H * DD + DD * tl.arange(0, H)[:, - None] + tl.arange(0, - D)[ - None, :]).to(tl.float32) + gq_ptr + pid * H * DD + DD * tl.arange(0, H)[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) gq_1 = tl.load( - gq_ptr + pid * H * DD + D + DD * tl.arange(0, H)[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) - - 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 + + pid * H * DD + + D + + DD * tl.arange(0, H)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) + + 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_0 = gq_0 * cos + gq_r * sin - q0 = tl.load(q_ptr + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q0 = tl.load( + q_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) q1 = tl.load( - q_ptr + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + q_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: s0 = tl.sigmoid(q0) @@ -2112,8 +2696,7 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -2127,8 +2710,7 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -2138,19 +2720,24 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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) - tl.store(dq_ptr + pid * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dq_1) + tl.store( + dq_ptr + pid * grad_stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + dq_0, + ) + tl.store( + dq_ptr + + pid * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dq_1, + ) 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) + 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 @@ -2170,27 +2757,31 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, dk_ptr = dqkv_ptr + DD * H gk_0 = tl.load( - gk_ptr + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :]).to(tl.float32) + gk_ptr + pid * h * DD + DD * tl.arange(0, h)[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) gk_1 = tl.load( - gk_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) - - 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 + + pid * h * DD + + D + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) + + 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_0 = gk_0 * cos + gk_r * sin - k0 = tl.load(k_ptr + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k0 = tl.load( + k_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) k1 = tl.load( - k_ptr + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + k_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: @@ -2199,8 +2790,7 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -2214,8 +2804,7 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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] + 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) @@ -2225,12 +2814,18 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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) - tl.store(dk_ptr + pid * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dk_1) + tl.store( + dk_ptr + pid * grad_stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + dk_0, + ) + tl.store( + dk_ptr + + pid * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dk_1, + ) 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) @@ -2246,23 +2841,23 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, dv_ptr = dqkv_ptr + DD * H + DD * h gv_0 = tl.load( - gv_ptr + pid * h * DD + DD * tl.arange(0, h)[:, - None] + tl.arange(0, - D)[ - None, :]).to(tl.float32) + gv_ptr + pid * h * DD + DD * tl.arange(0, h)[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) gv_1 = tl.load( - gv_ptr + pid * h * DD + D + DD * tl.arange(0, h)[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + gv_ptr + + pid * h * DD + + D + + DD * tl.arange(0, h)[:, None] + + tl.arange(0, D)[None, :] + ).to(tl.float32) if SILU: - v0 = tl.load(v_ptr + pid * stride + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v0 = tl.load( + v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) v1 = tl.load( - v_ptr + pid * stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :]).to(tl.float32) + v_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + ).to(tl.float32) s0 = tl.sigmoid(v0) s1 = tl.sigmoid(v1) @@ -2272,43 +2867,52 @@ def varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, gv_ptr, 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) - tl.store(dv_ptr + pid * grad_stride + D + DD * row_offs[:, - None] + tl.arange( - 0, D)[None, :], dv_1) + tl.store( + dv_ptr + pid * grad_stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + dv_0, + ) + tl.store( + dv_ptr + + pid * grad_stride + + D + + DD * row_offs[:, None] + + tl.arange(0, D)[None, :], + dv_1, + ) @triton.jit -def compatible_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 - ): +def compatible_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 @@ -2336,29 +2940,38 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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_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_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) + 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) + 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) @@ -2366,8 +2979,7 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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] + 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) @@ -2381,8 +2993,7 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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] + 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) @@ -2392,19 +3003,26 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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( + 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) + 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 @@ -2426,28 +3044,38 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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_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_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) + 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) + k_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, + ).to(tl.float32) if SILU: @@ -2456,8 +3084,7 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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] + 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) @@ -2471,8 +3098,7 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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] + 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) @@ -2482,12 +3108,20 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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( + 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) @@ -2505,24 +3139,34 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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_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) + 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) + 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) + 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) @@ -2532,24 +3176,40 @@ def compatible_varlen_qk_norm_and_half_rope_backward_kernel(gq_ptr, gk_ptr, 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): + 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: @@ -2595,14 +3255,18 @@ def triton_varlen_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, if PH == H and ph == h: varlen_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, cu_seqlens_q, cu_seqlens_kv, dqkv, - tmp_dqw, tmp_dkw, + tmp_dqw, + tmp_dkw, B, stride, grad_stride, @@ -2619,18 +3283,22 @@ def triton_varlen_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, cp_size, reuse, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) else: compatible_varlen_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, cu_seqlens_q, cu_seqlens_kv, dqkv, - tmp_dqw, tmp_dkw, + tmp_dqw, + tmp_dkw, B, stride, grad_stride, @@ -2649,7 +3317,7 @@ def triton_varlen_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, cp_size, reuse, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) dqw = tmp_dqw.sum(0) dkw = tmp_dkw.sum(0) @@ -2657,28 +3325,32 @@ def triton_varlen_qk_norm_and_half_rope_backward(gq, gk, gv, qkv, q_norm_weight, @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 - ): +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) + pos = _get_varlen_token_idx(cu_seqlens_q_ptr, pid, B, PB, cp_rank, cp_size) L = num_tokens bid = 0 else: @@ -2693,14 +3365,18 @@ def mla_rope_forward_kernel(q_ptr, kv_ptr, k_pos_emb_ptr, 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) + 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)) + 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))) @@ -2709,87 +3385,130 @@ def mla_rope_forward_kernel(q_ptr, kv_ptr, k_pos_emb_ptr, 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) + 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) + 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) + 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) + 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) + 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) + 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 = 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, :]) + 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, :]) + 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) + 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) + 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) + 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) + 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) + 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) + 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): +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: @@ -2841,8 +3560,9 @@ def triton_mla_rope_forward(q, kv, k_pos_emb, freqs, mscale=1.0, 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()=}" + 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 @@ -2868,7 +3588,7 @@ def triton_mla_rope_forward(q, kv, k_pos_emb, freqs, mscale=1.0, False if VARLEN else transpose, reuse, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) if VARLEN or not transpose: qo = q @@ -2876,28 +3596,31 @@ def triton_mla_rope_forward(q, kv, k_pos_emb, freqs, mscale=1.0, @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 - ): +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) + pos = _get_varlen_token_idx(cu_seqlens_q_ptr, pid, B, PB, cp_rank, cp_size) L = num_tokens bid = 0 else: @@ -2912,17 +3635,22 @@ def mla_rope_backward_kernel(q_ptr, k_ptr, v_ptr, freqs_ptr, 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) + 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 @@ -2937,19 +3665,35 @@ def mla_rope_backward_kernel(q_ptr, k_ptr, v_ptr, freqs_ptr, 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, :]) + 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 + 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) + 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) + 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)), @@ -2958,12 +3702,10 @@ def mla_rope_backward_kernel(q_ptr, k_ptr, v_ptr, freqs_ptr, # 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) + 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) + 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 @@ -2972,62 +3714,91 @@ def mla_rope_backward_kernel(q_ptr, k_ptr, v_ptr, freqs_ptr, 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) + 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) + 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) + 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) + 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, :]) + 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) + 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) + 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, :]) + 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) + 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) + 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): +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 @@ -3080,7 +3851,7 @@ def triton_mla_rope_backward(q_grad, k_grad, v_grad, freqs, mscale=1.0, False if VARLEN else transposed, reuse, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) if VARLEN or not transposed: dq = q_grad diff --git a/linghe/utils/scatter.py b/linghe/utils/scatter.py index 8f8ada6..501c489 100644 --- a/linghe/utils/scatter.py +++ b/linghe/utils/scatter.py @@ -11,9 +11,16 @@ @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: @@ -59,13 +68,16 @@ 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 @@ -73,6 +85,7 @@ def triton_aligned_scatter_add(x: torch.Tensor, # for deepep scatter_add # 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): pid = tl.program_id(axis=0) @@ -82,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, sem='relaxed') + tl.atomic_add(o_ptr + dst_idx * N + offs, x, sem="relaxed") @triton.jit @@ -111,8 +124,9 @@ def triton_scatter_add(x, outputs, indices): 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 = 512 T = triton.cdiv(M, sm) @@ -122,21 +136,14 @@ 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 @@ -144,54 +151,49 @@ def triton_scatter_add(x, outputs, indices): @triton.jit 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, + 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 ** 24, 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): load_mask = (index >= 0) & (tl.arange(0, N) < n) - sums += tl.load(grads_ptr + index * n + tl.arange(0, N), - mask=load_mask).to( + 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 ** 24, 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, mask=tl.arange(0, N) < n - ) + 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 @@ -209,15 +211,14 @@ def triton_unpermute_with_mask_map( N = triton.next_power_of_2(n) num_tokens, num_experts = row_id_map.shape # not transposed - output = torch.empty((num_tokens, n), dtype=grad.dtype, - device="cuda") + output = torch.empty((num_tokens, n), dtype=grad.dtype, device="cuda") PROB = probs is not None if PROB: assert probs.is_contiguous() - restore_probs = torch.zeros((num_tokens, num_experts), - dtype=probs.dtype, - device="cuda") + restore_probs = torch.zeros( + (num_tokens, num_experts), dtype=probs.dtype, device="cuda" + ) else: restore_probs = None @@ -236,6 +237,6 @@ def triton_unpermute_with_mask_map( num_experts, PROB, num_stages=4, - num_warps=4 + num_warps=4, ) return output, restore_probs diff --git a/linghe/utils/silu.py b/linghe/utils/silu.py index 7b67dd1..b79f744 100644 --- a/linghe/utils/silu.py +++ b/linghe/utils/silu.py @@ -24,31 +24,28 @@ def exp2(x): @triton.jit def weighted_silu_forward_asm_kernel( - x_ptr, - weight_ptr, - out_ptr, - M, - N, - H: tl.constexpr, - W: tl.constexpr, - WEIGHT: tl.constexpr, + 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 offs = ( - rid * H * N + cid * W + tl.arange(0, H)[:, None] * N + tl.arange(0, - W)[ - None, :] + 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] + 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)) @@ -59,40 +56,41 @@ def weighted_silu_forward_asm_kernel( 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, :] + rid * H * n + cid * W + tl.arange(0, H)[:, None] * n + tl.arange(0, W)[None, :] ) tl.store(out_ptr + offs, x, mask=mask) @triton.jit -def weighted_silu_forward_kernel(x_ptr, weight_ptr, out_ptr, M, - N, - H: tl.constexpr, - W: tl.constexpr, - WEIGHT: tl.constexpr): +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, :] + 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) + 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] + 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, :] + 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) @@ -120,47 +118,33 @@ def triton_weighted_silu_forward(x, weight=None, out=None, asm=False): 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 + 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 + 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, - N, - H: 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 * 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, :] + 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] @@ -190,9 +174,9 @@ def weighted_silu_backward_kernel(g_ptr, x_ptr, weight_ptr, dx_ptr, dw_ptr, 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: @@ -221,39 +205,34 @@ def triton_weighted_silu_backward(g: torch.Tensor, 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 + 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, - M, - n: tl.constexpr, - H: tl.constexpr, - W: tl.constexpr, - ROUND: tl.constexpr, - OUTPUT_MODE: tl.constexpr): +def silu_and_block_quant_forward_kernel( + x_ptr, + out_ptr, + scale_ptr, + transpose_output_ptr, + transpose_scale_ptr, + M, + n: tl.constexpr, + H: tl.constexpr, + W: 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, :] + 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 @@ -266,37 +245,40 @@ def silu_and_block_quant_forward_kernel(x_ptr, 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) + 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) + 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: 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) + tl.store(transpose_scale_ptr + rid * n + cid * W + tl.arange(0, W), scale) xq = (x / scale).to(transpose_output_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, - round_scale=False, - output_mode=2): + 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, round_scale=False, output_mode=2 +): """ fused silu and blockwise quantization, used in shared expert Args: @@ -320,13 +302,12 @@ def triton_silu_and_block_quant_forward(x, 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) + 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) + 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, 4 elif output_mode == 1: @@ -348,88 +329,99 @@ def triton_silu_and_block_quant_forward(x, round_scale, output_mode, num_stages=2, - num_warps=num_warps + 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, - M, - n: tl.constexpr, - ROUND: tl.constexpr): +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, +): 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, :] + 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) + 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)) - scale1 = tl.maximum( - tl.max(dx1.abs(), 1) / 448, 1e-30) + dx1 = sigmoid * g * x2 * (1 + x1 * (1 - sigmoid)) + 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) + 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) + 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) + 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) dx2 = sigmoid * g * x1 - scale2 = tl.maximum( - tl.max(dx2.abs(), 1) / 448, 1e-30) + 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) + 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) + 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) + 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) + 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, - round_scale=False): +def triton_silu_and_block_quant_backward(g, x, round_scale=False): """ backward of triton_silu_and_block_quant_forward Args: @@ -452,8 +444,7 @@ def triton_silu_and_block_quant_backward(g, x, 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) + transpose_dx_scale = torch.empty(scale_shape, device=device, dtype=torch.float32) assert M % 128 == 0 and N % 256 == 0 grid = (M // 128, N // 256) @@ -468,22 +459,25 @@ def triton_silu_and_block_quant_backward(g, x, n, round_scale, num_stages=2, - num_warps=8 + 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, - n, - E: tl.constexpr, - ROUND: tl.constexpr): +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, + E: tl.constexpr, + ROUND: tl.constexpr, +): eid = tl.program_id(axis=0) rid = tl.program_id(axis=1) cid = tl.program_id(axis=2) @@ -503,23 +497,32 @@ def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, 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, :] + 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) 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] @@ -528,7 +531,9 @@ def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, scale = tl.exp2(tl.ceil(tl.log2(scale))) tl.store( scale_ptr + si * nb + cid * count + rid * 128 + tl.arange(0, 128), - scale, mask=indices < count) + scale, + mask=indices < count, + ) xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) tl.store(out_ptr + hoffs, xq, mask=mask) @@ -537,26 +542,33 @@ def batch_weighted_silu_and_block_quant_forward_kernel(x_ptr, weight_ptr, if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) tl.store( - transpose_scale_ptr + transpose_scale_off * n + rid * n + cid * 128 + tl.arange( - 0, 128), scale) + 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) + tl.store(transpose_output_ptr + toffs, xq, 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, - n, - B: tl.constexpr, - E: tl.constexpr, - ROUND: tl.constexpr): +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, + n, + B: 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) @@ -573,30 +585,34 @@ def batch_weighted_silu_and_block_quant_forward_nt_kernel(x_ptr, weight_ptr, 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, :] + 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) for i in range(I): 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] 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) + tl.store(scale_ptr + soffs, scale, mask=indices < count) xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) @@ -610,51 +626,56 @@ def batch_weighted_silu_and_block_quant_forward_nt_kernel(x_ptr, weight_ptr, counts = tl.load(count_ptr + tl.arange(0, E)) n_blocks = tl.cdiv(counts, 128) transpose_soff = 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, B)[None, :] - toffs = si * n + rid * 128 + cid * count * 128 + tl.arange(0, B)[:, - None] * count + tl.arange( - 0, 128)[ - None, :] - tsoffs = transpose_soff * n + rid * n + cid * 128 + tl.arange( - 0, B) + offs = ( + si * n * 2 + + rid * 128 * n * 2 + + cid * 128 + + tl.arange(0, 128)[:, None] * n * 2 + + tl.arange(0, B)[None, :] + ) + toffs = ( + si * n + + rid * 128 + + cid * count * 128 + + tl.arange(0, B)[:, None] * count + + tl.arange(0, 128)[None, :] + ) + tsoffs = transpose_soff * n + rid * n + cid * 128 + tl.arange(0, B) indices = rid * 128 + tl.arange(0, 128) for i in range(I): 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] 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 + tsoffs, scale) + tl.store(transpose_scale_ptr + tsoffs, scale) xq = tl.trans((x / scale).to(transpose_output_ptr.dtype.element_ty)) - tl.store(transpose_output_ptr + toffs, xq, - mask=indices[None, :] < count) + tl.store(transpose_output_ptr + toffs, xq, mask=indices[None, :] < count) offs += B toffs += count * B tsoffs += 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, - n, - B: tl.constexpr, - E: tl.constexpr, - ROUND: tl.constexpr): +def batch_weighted_silu_and_block_quant_forward_n_kernel( + x_ptr, + weight_ptr, + out_ptr, + scale_ptr, + count_ptr, + accum_ptr, + n, + B: 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) @@ -669,16 +690,18 @@ def batch_weighted_silu_and_block_quant_forward_n_kernel(x_ptr, weight_ptr, 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, :] + 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) + 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] @@ -687,25 +710,34 @@ def batch_weighted_silu_and_block_quant_forward_n_kernel(x_ptr, weight_ptr, 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) + 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, :] + 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, - n, - B: tl.constexpr, - E: tl.constexpr, - ROUND: tl.constexpr): +def batch_weighted_silu_and_block_quant_forward_t_kernel( + x_ptr, + weight_ptr, + transpose_output_ptr, + transpose_scale_ptr, + count_ptr, + accum_ptr, + n, + B: 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) @@ -723,16 +755,18 @@ def batch_weighted_silu_and_block_quant_forward_t_kernel(x_ptr, weight_ptr, 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, :] + 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) + 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] @@ -740,27 +774,36 @@ def batch_weighted_silu_and_block_quant_forward_t_kernel(x_ptr, weight_ptr, 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) + 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, - round_scale=False, - output_mode=2): + 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, + round_scale=False, + output_mode=2, +): """ silu and blockwise quantize activation in routed experts Args: @@ -780,7 +823,7 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, - 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 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 @@ -793,10 +836,8 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, 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) + 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 @@ -818,7 +859,7 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, len(splits), round_scale, num_stages=2, - num_warps=2 + num_warps=2, ) elif output_mode == 1: B = 32 @@ -835,7 +876,7 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, len(splits), round_scale, num_stages=2, - num_warps=2 + num_warps=2, ) else: grid = (n_experts, triton.cdiv(max(splits), 128), n // 128) @@ -852,7 +893,7 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, len(splits), round_scale, num_stages=2, - num_warps=8 + num_warps=8, ) # B = 16 @@ -910,25 +951,28 @@ def triton_batch_weighted_silu_and_block_quant_forward(x, @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, - E: tl.constexpr, - ROUND: tl.constexpr): +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, + 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) si = tl.load(accum_ptr + eid) - count - # very slow with triton 3.3.1, fix in 3.5.1 + # 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)) @@ -937,12 +981,19 @@ def batch_weighted_silu_and_block_quant_backward_kernel(g_ptr, x_ptr, 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)) + 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, :] + 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) @@ -950,78 +1001,114 @@ def batch_weighted_silu_and_block_quant_backward_kernel(g_ptr, x_ptr, 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) + 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) dw = tl.sum(sigmoid * x1 * 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) + 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) + tl.store( + dx_scale_ptr + si * nb * 2 + cid * count + rid * 128 + tl.arange(0, 128), + scale, + mask=idx < count, + ) tl.store(dx_ptr + offs, dx / scale[:, None], mask=idx[:, None] < count) - scale = tl.maximum( - tl.max(dx.abs(), 0) / 448, 1e-30) + 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 * 128 + tl.arange( - 0, 128), scale) + transpose_dx_scale_ptr + + transpose_off * n * 2 + + rid * n * 2 + + cid * 128 + + tl.arange(0, 128), + scale, + ) qdx = tl.trans((dx / scale[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) + 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 = sigmoid * g * x1 * w - scale = tl.maximum( - tl.max(dx.abs(), 1) / 448, 1e-30) + 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 + count * nb + tl.arange( - 0, 128), scale, mask=idx < count) + dx_scale_ptr + + si * nb * 2 + + cid * count + + rid * 128 + + count * nb + + tl.arange(0, 128), + scale, + mask=idx < count, + ) tl.store(dx_ptr + n + offs, dx / scale[:, None], mask=idx[:, None] < count) - scale = tl.maximum( - tl.max(dx.abs(), 0) / 448, 1e-30) + 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)) 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 + + transpose_off * n * 2 + + rid * n * 2 + + n + + cid * 128 + + tl.arange(0, 128), + scale, + ) 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 + + 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, - n, - B: tl.constexpr, - E: tl.constexpr, - ROUND: tl.constexpr): +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, + n, + B: 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) @@ -1037,56 +1124,75 @@ def batch_weighted_silu_and_block_quant_backward_n_kernel(g_ptr, x_ptr, 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, :] + 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) - 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) + 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) sigmoid = tl.sigmoid(x1) dw = tl.sum(sigmoid * x1 * 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) + 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_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 = sigmoid * g * x1 * w - scale = tl.maximum( - tl.max(dx.abs(), 1) / 448, 1e-30) + 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) + 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, - n, - B: tl.constexpr, - E: tl.constexpr, - ROUND: tl.constexpr): +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, + n, + B: 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) @@ -1100,61 +1206,93 @@ def batch_weighted_silu_and_block_quant_backward_t_kernel(g_ptr, x_ptr, 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)) + 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, :] + 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) - 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) + 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) sigmoid = tl.sigmoid(x1) dx = sigmoid * g * x2 * w * (1 + x1 * (1 - sigmoid)) - scale = tl.maximum( - tl.max(dx.abs(), 0) / 448, 1e-30) + 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) + 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) + 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 = sigmoid * g * x1 * w - scale = tl.maximum( - tl.max(dx.abs(), 0) / 448, 1e-30) + 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) + 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) + 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, - round_scale=False): +def triton_batch_weighted_silu_and_block_quant_backward( + g, x, weight, counts, splits=None, round_scale=False +): """ backward of triton_batch_weighted_silu_and_block_quant_forward Args: @@ -1176,7 +1314,7 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, 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' + assert splits is not None, "batch mode need splits to launch kernels" device = x.device @@ -1186,10 +1324,8 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, 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) + 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 @@ -1232,7 +1368,7 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, n_expert, round_scale, num_stages=2, - num_warps=4 + num_warps=4, ) dw = dws.sum(1, keepdim=True).to(weight.dtype) @@ -1251,20 +1387,27 @@ def triton_batch_weighted_silu_and_block_quant_backward(g, x, weight, n_expert, round_scale, num_stages=2, - num_warps=4 + num_warps=4, ) 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, + max_ptr, + M, + T, + n: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, + CALIBRATE: tl.constexpr, +): pid = tl.program_id(axis=0) row_offs = pid * T * W * n + tl.arange(0, W)[:, None] * n @@ -1278,8 +1421,7 @@ def silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out_ptr, 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) + x2 = tl.load(x_ptr + n + row_offs * 2 + col_offs, mask=mask).to(tl.float32) x = x1 * tl.sigmoid(x1) * x2 if CALIBRATE: maxs = tl.maximum(x.abs(), maxs) @@ -1299,14 +1441,19 @@ def silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, out_ptr, # 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, + max_ptr, + M, + T: tl.constexpr, + n: tl.constexpr, + B: tl.constexpr, + ROUND: tl.constexpr, + CALIBRATE: tl.constexpr, +): pid = tl.program_id(axis=0) # rowwise read with block size [T, B] @@ -1348,10 +1495,15 @@ def compatible_silu_and_smooth_quant_forward_kernel(x_ptr, smooth_scale_ptr, # 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=None, + out=None, + scale=None, + maxs=None, + round_scale=False, + calibrate=False, +): """""" assert x.is_contiguous() M, N = x.shape @@ -1384,7 +1536,7 @@ def triton_silu_and_smooth_quant_forward(x, smooth_scale=None, out=None, round_scale, calibrate, num_stages=2, - num_warps=16 + num_warps=16, ) else: B = 512 @@ -1406,7 +1558,7 @@ def triton_silu_and_smooth_quant_forward(x, smooth_scale=None, out=None, round_scale, calibrate, num_stages=2, - num_warps=16 + num_warps=16, ) if calibrate: @@ -1416,29 +1568,32 @@ def triton_silu_and_smooth_quant_forward(x, smooth_scale=None, out=None, @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] + 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)) @@ -1456,8 +1611,7 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, # g = tl.load(g_ptr + hoffs) # 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 @@ -1465,29 +1619,31 @@ def 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(transpose_dx_ptr.dtype.element_ty))) - tl.store(transpose_dx_scale_ptr + pid * n * 2 + i * B + tl.arange(0, B), - 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) + 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 @@ -1501,10 +1657,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)) @@ -1516,8 +1670,7 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, x2 = tl.load(x_ptr + offs + n).to(tl.float32) g = tl.load(g_ptr + hoffs).to(tl.float32) sigmoid = tl.sigmoid(x1) - dx1 = g * x2 * sigmoid * ( - 1 + x1 * (1 - sigmoid)) * smooth_scale_1 + 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) @@ -1531,17 +1684,14 @@ def silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, # 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)) @@ -1552,11 +1702,14 @@ def _requant_kernel(x_ptr, scale_ptr, scales_ptr, # 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 @@ -1567,14 +1720,12 @@ def triton_silu_and_smooth_quant_backward(g, x, dx_scale = torch.empty((M,), device=device, dtype=torch.float32) 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) + transpose_dx_scale = torch.empty(scale_shape, device=device, dtype=torch.float32) T = 32 B = 32 assert M % T == 0 and n % B == 0 - transpose_dx_scales = torch.empty((M // T, N), device=device, - dtype=torch.float32) + transpose_dx_scales = torch.empty((M // T, N), device=device, dtype=torch.float32) grid = (M // T,) silu_and_smooth_quant_backward_kernel[grid]( g, @@ -1592,32 +1743,34 @@ 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, + max_ptr, + count_ptr, + accum_ptr, + M, + n: tl.constexpr, + W: tl.constexpr, + ROUND: tl.constexpr, + REVERSE: tl.constexpr, + CALIBRATE: tl.constexpr, +): eid = tl.program_id(axis=0) tid = tl.program_id(axis=1) sm = tl.num_programs(axis=1) @@ -1640,12 +1793,11 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, 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] + w = tl.load(weight_ptr + si + indices, mask=indices < count).to(tl.float32)[ + :, None + ] x = x1 * tl.sigmoid(x1) * x2 if CALIBRATE: @@ -1666,16 +1818,18 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, # 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, + calibrate=False, +): """""" assert x.is_contiguous() and weight.is_contiguous() M, N = x.shape @@ -1691,14 +1845,11 @@ def triton_batch_weighted_silu_and_smooth_quant_forward(x, 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) + 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) + 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 @@ -1724,7 +1875,7 @@ def triton_batch_weighted_silu_and_smooth_quant_forward(x, reverse, calibrate, num_stages=3, - num_warps=16 + num_warps=16, ) if calibrate: maxs = tmp_maxs.amax(1) @@ -1733,23 +1884,26 @@ def triton_batch_weighted_silu_and_smooth_quant_forward(x, @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) @@ -1761,52 +1915,71 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel(g_ptr, x_ptr, 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, :] + 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, :] + ) 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 = tl.sigmoid(x1) - dx1 = g * x2 * sigmoid * ( - 1 + x1 * (1 - sigmoid)) * w + dx1 = g * x2 * sigmoid * (1 + x1 * (1 - sigmoid)) * w dx2 = g * x1 * sigmoid * w dw += tl.sum(x1 * sigmoid * x2 * g, 1) @@ -1816,27 +1989,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 @@ -1846,33 +2035,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 = tl.sigmoid(x1) - dx1 = g * x2 * sigmoid * ( - 1 + x1 * (1 - sigmoid)) * smooth_scale_1 * w + 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) @@ -1886,13 +2076,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) @@ -1903,17 +2096,26 @@ 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)) + 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, :] + 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] @@ -1921,13 +2123,17 @@ 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 @@ -1935,7 +2141,7 @@ def triton_batch_weighted_silu_and_smooth_quant_backward(g, x, weight, 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 @@ -1951,16 +2157,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]( @@ -1983,18 +2190,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 index 8848b71..1c7e782 100644 --- a/linghe/utils/topk.py +++ b/linghe/utils/topk.py @@ -9,9 +9,9 @@ @triton.jit -def topk_forward_kernel(input_ptr, value_ptr, index_ptr, - N: tl.constexpr, - K: tl.constexpr): +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)) @@ -41,7 +41,7 @@ def triton_topk_forward(x, k, dim=-1): x: input tensor. k: topk Returns: - values: topk values + values: topk values indices: topk indices """ device = x.device @@ -59,22 +59,12 @@ def triton_topk_forward(x, k, dim=-1): values = torch.empty((M, k), device=device, dtype=x.dtype) indices = torch.empty((M, k), device=device, dtype=torch.int64) grid = (g,) - topk_forward_kernel[grid]( - x, - values, - indices, - N, - k, - num_stages=2, - num_warps=2 - ) + topk_forward_kernel[grid](x, values, indices, N, k, num_stages=2, num_warps=2) return values, indices @triton.jit -def topk_backward_kernel(grad_ptr, index_ptr, dx_ptr, - N: tl.constexpr, - K: tl.constexpr): +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)) @@ -106,27 +96,25 @@ def triton_topk_backward(grad_output, indices, N, dim=-1): dx = torch.zeros((M, 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 + grad_output, indices, dx, N, k, num_stages=2, num_warps=2 ) return dx @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 - ): +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 @@ -157,8 +145,7 @@ def group_topk_score_forward_kernel(input_ptr, bias_ptr, prob_ptr, map_ptr, 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 + 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) @@ -171,27 +158,31 @@ def group_topk_score_forward_kernel(input_ptr, bias_ptr, prob_ptr, map_ptr, 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)) + 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)) + 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) -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): +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: @@ -199,13 +190,13 @@ def triton_group_topk_score_forward(x, k, expert_bias: expert bias k: topk Returns: - probs: - routing_map: - tokens_per_expert: + probs: + routing_map: + tokens_per_expert: """ device = x.device shape = x.shape - assert len(shape) <= 3 and x.is_contiguous() and score_function == 'sigmoid' + assert len(shape) <= 3 and x.is_contiguous() and score_function == "sigmoid" if len(shape) == 3: M, B, N = shape g = M * B @@ -231,16 +222,15 @@ def triton_group_topk_score_forward(x, k, group_topk, BIAS, num_stages=1, - num_warps=1 + 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): +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)) @@ -252,8 +242,9 @@ def group_topk_score_backward_kernel(grad_ptr, input_ptr, map_ptr, dx_ptr, 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): +def triton_group_topk_score_backward( + grad_output, input, routing_map, scaling_factor=1.0, eps=1e-20 +): """ topk backward. Args: @@ -264,8 +255,9 @@ def triton_group_topk_score_backward(grad_output, input, routing_map, """ device = grad_output.device shape = grad_output.shape - assert len( - shape) <= 3 and grad_output.is_contiguous() and routing_map.is_contiguous() + 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 @@ -284,6 +276,6 @@ def triton_group_topk_score_backward(grad_output, input, routing_map, eps, N, num_stages=2, - num_warps=1 + num_warps=1, ) return dx diff --git a/linghe/utils/transpose.py b/linghe/utils/transpose.py index b99ad35..ae65cb6 100644 --- a/linghe/utils/transpose.py +++ b/linghe/utils/transpose.py @@ -12,37 +12,44 @@ from linghe.tools.util import round_up - # os.environ["TRITON_PRINT_AUTOTUNING"] = "1" @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)) + 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), + ) @triton.jit -def transpose_inner_dims_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) @@ -52,31 +59,43 @@ def transpose_inner_dims_kernel(x_ptr, t_ptr, B, M, b_stride, m_stride, @triton.jit -def transpose_outer_dims_kernel(x_ptr, t_ptr, M, N, H: tl.constexpr, - W: tl.constexpr, - EVEN: tl.constexpr): +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, :] + 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: - 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)) + 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), + ) def triton_transpose(x: torch.Tensor, inner=True): @@ -84,7 +103,7 @@ def triton_transpose(x: torch.Tensor, inner=True): transpose x with dim0 and dim1 Args: x: input tensor - inner: inner dim if True, outer dim if False + inner: inner dim if True, outer dim if False Returns: transposed tensor @@ -104,20 +123,14 @@ def triton_transpose(x: torch.Tensor, inner=True): 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 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) @@ -126,22 +139,22 @@ def triton_transpose(x: torch.Tensor, inner=True): num_stages = 5 num_warps = 2 grid = (B, M) - transpose_inner_dims_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: 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) + 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) @@ -154,42 +167,33 @@ def triton_transpose(x: torch.Tensor, inner=True): 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 + 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 transpose_and_pad_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): @@ -220,25 +224,23 @@ def triton_transpose_and_pad(x, out=None, pad=True): 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 + 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) @@ -259,11 +261,10 @@ def triton_batch_transpose(xs, xts=None): n_experts = len(xs) device = xs[0].device if xts is None: - 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) + 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 @@ -271,20 +272,23 @@ def triton_batch_transpose(xs, xts=None): 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, [M] * n_experts) 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_transpose_and_pad_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) @@ -292,13 +296,15 @@ 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 @@ -326,8 +332,10 @@ def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): pad_sizes = [round_up(x, b=32) 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: @@ -338,13 +346,16 @@ def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): num_warps = 8 grid = (n_experts, N // W) batch_transpose_and_pad_kernel[grid]( - x, x_t, - counts, accums, + 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) @@ -365,13 +376,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)) @@ -387,8 +396,5 @@ def triton_opt_transpose(x): 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 14343d7..eb9ce22 100644 --- a/linghe/utils/unary.py +++ b/linghe/utils/unary.py @@ -9,11 +9,16 @@ @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: @@ -29,8 +34,9 @@ def calculate_smooth_scale_kernel(x_ptr, y_ptr, min_value, smooth_coef, tl.store(y_ptr + offs, x, mask=offs < N) -def triton_calculate_smooth_scale(x, min_value=1.0, smooth_coef=0.5, - inplace=False, round_scale=False): +def triton_calculate_smooth_scale( + x, min_value=1.0, smooth_coef=0.5, inplace=False, round_scale=False +): assert x.is_contiguous() N = x.shape[0] B = 4096 @@ -46,7 +52,8 @@ def triton_calculate_smooth_scale(x, min_value=1.0, smooth_coef=0.5, num_warps = 4 grid = (triton.cdiv(N, B),) calculate_smooth_scale_kernel[grid]( - x, output, + x, + output, min_value, smooth_coef, N, @@ -54,15 +61,15 @@ def triton_calculate_smooth_scale(x, min_value=1.0, smooth_coef=0.5, EVEN, round_scale, num_stages=num_stages, - num_warps=num_warps + num_warps=num_warps, ) return output @triton.jit -def batch_clip_kernel(input_ptrs, size_ptr, clip_value, - DT: tl.constexpr, - B: tl.constexpr): +def batch_clip_kernel( + input_ptrs, size_ptr, clip_value, DT: tl.constexpr, B: tl.constexpr +): tid = tl.program_id(axis=0) bid = tl.program_id(axis=1) T = tl.num_programs(axis=1) @@ -77,8 +84,7 @@ def batch_clip_kernel(input_ptrs, size_ptr, clip_value, 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)) + tl.store(input_ptr + offs, xc, mask=(offs < size) & (tl.abs(x) > clip_value)) offs += B @@ -99,23 +105,17 @@ def triton_batch_clip(xs, clip_value=100.0): 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) + 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 grid = (tensor_count, T) - batch_clip_kernel[grid]( - ptrs, - sizes, - clip_value, - DT, - B, - num_stages=2, - num_warps=2 - ) + 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 1cc5045..f8a7b2e 100644 --- a/scripts/plot_input_output.py +++ b/scripts/plot_input_output.py @@ -2,22 +2,22 @@ import torch -def read_bf16_inputs(prefix='fc2'): +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() @@ -27,31 +27,31 @@ def read_bf16_inputs(prefix='fc2'): return x, w, dy, dx, dw -def read_fp8_inputs(prefix='fc2'): +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 + 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'] + 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'] + 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'] + 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'] + dytq = d["dyt"].float().t() + dyts = d["dyts"] + dytm = d["dyt_smooth_scale"] dyt = dytq / dytm * dyts[:, None] x = x.cuda() @@ -65,7 +65,7 @@ def read_fp8_inputs(prefix='fc2'): # bf16 if True: - prefix = 'out' + prefix = "out" x, w, dy, dx, dw = read_bf16_inputs(prefix=prefix) r = 5 @@ -75,45 +75,45 @@ def read_fp8_inputs(prefix='fc2'): 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' + prefix = "fc2" x, w, dy, dyt, xm, wm = read_fp8_inputs(prefix=prefix) r = 5 @@ -121,24 +121,24 @@ def read_fp8_inputs(prefix='fc2'): 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 d8c192a..ff567b0 100644 --- a/scripts/reproduce_triton_bug.py +++ b/scripts/reproduce_triton_bug.py @@ -9,27 +9,31 @@ @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, :] + 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) @@ -54,10 +58,9 @@ def rms_norm_and_smooth_quant_forward_kernel(x_ptr, weight_ptr, # 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, - num_warps=4): +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 @@ -92,37 +95,44 @@ def triton_rms_norm_and_quant_forward(x, weight, smooth_scale, eps=1e-6, calibrate, round_scale, num_stages=3, - num_warps=num_warps + num_warps=num_warps, ) if calibrate: maxs = maxs.amax(0) return out, scale, maxs, rms -if __name__ == '__main__': +if __name__ == "__main__": M = 4096 N = 2048 dtype = torch.bfloat16 - device = 'cuda:0' + 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 + ) # 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=}') + 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=}') + 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/setup.py b/setup.py index ae32ca6..5e1df95 100644 --- a/setup.py +++ b/setup.py @@ -9,8 +9,9 @@ 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", diff --git a/tests/test_add.py b/tests/test_add.py index b56fb8e..e6babd7 100644 --- a/tests/test_add.py +++ b/tests/test_add.py @@ -20,7 +20,7 @@ def torch_add(x, outputs, accum=True): def test_triton_inplace_add(M=4096, N=4096, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" outputs = torch.randn(M, N, dtype=dtype, device=device) x = torch.randn(M, N, dtype=dtype, device=device) @@ -28,22 +28,32 @@ def test_triton_inplace_add(M=4096, N=4096, bench=False): out = outputs.clone() triton_inplace_add(out, x) out_ref = outputs + x - output_check(out_ref, out, 'sum') + output_check(out_ref, out, "sum") if bench: n_repeat = 100 - 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) - - 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) - - -if __name__ == '__main__': + 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, + ) + + 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, + ) + + +if __name__ == "__main__": test_triton_inplace_add(M=4096, N=4096) diff --git a/tests/test_blockwise_fp8_gemm.py b/tests/test_blockwise_fp8_gemm.py index 9e5d3e6..611cf72 100644 --- a/tests/test_blockwise_fp8_gemm.py +++ b/tests/test_blockwise_fp8_gemm.py @@ -5,15 +5,14 @@ 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_bb_fp8_gemm, triton_tt_fp8_gemm from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check def test_triton_bb_gemm(M=4096, N=4096, K=4096, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" B = 64 x = torch.randn(M, K, dtype=dtype, device=device) @@ -25,28 +24,37 @@ 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) + 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.to(dtype), y, name='y', rtol=0.05, atol=1.0) + y = triton_bb_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) 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) + 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' + device = "cuda:0" B = 64 x = torch.randn(M, K, dtype=dtype, device=device) @@ -62,19 +70,26 @@ def test_triton_tt_gemm(M=4096, N=4096, K=4096, bench=False): 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.to(dtype), y, 'y', atol=1.0, rtol=0.05) + y = triton_tt_fp8_gemm(x_q, w_q, x_scales, w_scales, out_dtype=dtype, block_size=B) + output_check(y_ref.to(dtype), y, "y", atol=1.0, rtol=0.05) 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) + 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__': +if __name__ == "__main__": test_triton_bb_gemm(M=4096, N=8192, K=2048, bench=False) test_triton_tt_gemm(M=4096, N=8192, K=2048, bench=False) diff --git a/tests/test_blockwise_quant.py b/tests/test_blockwise_quant.py index e921e34..0e726ef 100644 --- a/tests/test_blockwise_quant.py +++ b/tests/test_blockwise_quant.py @@ -1,17 +1,20 @@ import torch -from linghe.quant.block import triton_block_quant, triton_blockwise_quant, \ - triton_batch_blockwise_quant +from linghe.quant.block import ( + triton_block_quant, + triton_blockwise_quant, + triton_batch_blockwise_quant, +) from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check -from linghe.tools.util import (torch_block_quant, - torch_blockwise_quant, - torch_make_indices) +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): +def torch_batch_blockwise_quant(x, token_count_per_expert_list, round_scale=True): M, DIM = x.shape q_refs = [] s_refs = [] @@ -22,12 +25,12 @@ def torch_batch_blockwise_quant(x, c = token_count_per_expert_list[i] if c == 0: continue - y = x[s:s + c] + 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) + 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)) @@ -41,75 +44,72 @@ def torch_batch_blockwise_quant(x, def test_block_quant(M=8192, N=4096, bench=False): - device = 'cuda:0' - x = torch.randn((M, N), dtype=torch.bfloat16, - device=device) ** 3 + 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') + output_check(x_q_ref.float(), x_q.float(), "data") + output_check(x_s_ref.float(), x_s.float(), "scale") if bench: - benchmark_func(triton_block_quant, x, - round_scale=True, - ref_bytes=M * N * 4) + benchmark_func(triton_block_quant, x, round_scale=True, ref_bytes=M * N * 4) def test_blockwise_quant(M=8192, N=4096, bench=False): - device = 'cuda:0' - x = torch.randn((M, N), dtype=torch.bfloat16, - device=device) ** 3 + 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_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') + 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") if bench: - benchmark_func(triton_blockwise_quant, x, - round_scale=True, - ref_bytes=M * N * 4) + benchmark_func(triton_blockwise_quant, x, round_scale=True, ref_bytes=M * N * 4) def test_batch_block_quant(M=16384, N=2048, n_experts=32, topk=2, bench=False): - device = 'cuda:0' - logits = torch.randn((M, n_experts), dtype=torch.float32, - device=device) ** 3 + 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) + 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_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') + 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") if bench: - benchmark_func(triton_batch_blockwise_quant, x, token_count_per_expert, - token_count_per_expert_list, - round_scale=True, - ref_bytes=M * N * 4) + benchmark_func( + triton_batch_blockwise_quant, + x, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=M * N * 4, + ) -if __name__ == '__main__': +if __name__ == "__main__": test_block_quant(M=8192, N=4096, bench=False) test_blockwise_quant(M=8192, N=4096, bench=False) test_batch_block_quant(M=16384, N=2048, n_experts=32, topk=2, bench=False) diff --git a/tests/test_channel_quant.py b/tests/test_channel_quant.py index f08a64b..96943c3 100644 --- a/tests/test_channel_quant.py +++ b/tests/test_channel_quant.py @@ -5,41 +5,53 @@ import torch -from linghe.quant.channel import (triton_deprecated_tokenwise_row_quant, - triton_row_quant, - triton_tokenwise_row_quant) +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.check import output_check from linghe.tools.util import torch_row_quant def test_row_quant(M=4096, N=4096, round_scale=True, bench=False): - device = 'cuda:0' + 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, x_q, name='data') - output_check(x_scale_ref, x_scale, name='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, x_q, name='data') - output_check(x_scale_ref, x_scale, name='scale') + output_check(x_q_ref, x_q, name="data") + output_check(x_scale_ref, x_scale, name="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__': + 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) diff --git a/tests/test_channelwise_fp8_gemm.py b/tests/test_channelwise_fp8_gemm.py index b5d5490..7bb2b0d 100644 --- a/tests/test_channelwise_fp8_gemm.py +++ b/tests/test_channelwise_fp8_gemm.py @@ -12,12 +12,14 @@ 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) @@ -28,7 +30,7 @@ def scaled_gemm_and_update(x_q, w_q, x_scales, w_scales, c=None, accum=False): def test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False): 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) @@ -38,12 +40,10 @@ 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) + 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) + output_check(y_ref, y, name="y", atol=-1) if bench: y_bf16 = torch.randn(M, N, dtype=dtype, device=device) @@ -53,27 +53,86 @@ def test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False): 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__': + 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=False) diff --git a/tests/test_dist_loss.py b/tests/test_dist_loss.py index b825c01..52afe0f 100644 --- a/tests/test_dist_loss.py +++ b/tests/test_dist_loss.py @@ -2,6 +2,7 @@ """ Copyright (c) Ant Financial Service Group and its affiliates. """ + import os import random from datetime import timedelta @@ -11,37 +12,46 @@ 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 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'): + +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) + 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): +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_module.set_device(torch.device(f"cuda:{group_rank}")) - device = 'cuda' + device = "cuda" dtype = torch.bfloat16 - local_logits = torch.randn((M, N), dtype=dtype, device=device, - requires_grad=False) + local_logits = torch.randn((M, N), dtype=dtype, device=device, requires_grad=False) select = True if select: @@ -52,8 +62,9 @@ def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, grad_coef=1.0, targets = torch.tensor(targets, dtype=torch.long, device=device) else: - targets = torch.randint(0, N * group_size, (M,), dtype=torch.long, - device=device) + targets = torch.randint( + 0, N * group_size, (M,), dtype=torch.long, device=device + ) if ignore_index is not None: targets[:10] = ignore_index @@ -65,95 +76,145 @@ def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, grad_coef=1.0, 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_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) + 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 = 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) + 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) + 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, + local_logits.detach().clone(), + global_targets, + sum_exp, max_logit, global_output_grad, group, ignore_index=ignore_index, - inplace=inplace) + inplace=inplace, + ) - output_check(grad_ref[:, group_rank * N:(group_rank + 1) * N], grad, - name=f'grad:{group_rank}', digest=10) + 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 + # 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__': + 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)) + 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=False) - 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) + test_triton_softmax_cross_entropy( + M=8192, N=157184, coef=1.0, grad_coef=1.0, inplace=False, group=pg, bench=False + ) + 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_embedding.py b/tests/test_embedding.py index 64ab114..be9818c 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -5,32 +5,34 @@ import torch -from linghe.facade.emb import (embedding_lookup, - fused_accumulation_embedding_lookup) +from linghe.facade.emb import embedding_lookup, fused_accumulation_embedding_lookup from linghe.tools.benchmark import benchmark_func 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 - ) +from linghe.utils.emb import ( + triton_embedding_forward, + triton_embedding_backward, + triton_scan_and_count, + triton_sync_embedding_backward, + triton_atomic_embedding_backward, +) def test_scan(M=4096, bench=False): - device = 'cuda:0' + device = "cuda: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) + 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) + 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') + output_check(accum_counts_ref, accum_counts[:size], name="accum_counts") if bench: ref_time = benchmark_func(triton_scan_and_count, sorted_ids) @@ -38,11 +40,10 @@ def test_scan(M=4096, bench=False): def test_embedding(B=2, M=4096, V=150000, D=4096, transpose=False, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + 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) + input_ids = torch.randint(0, V // 15, (B, M), dtype=torch.int32, device=device) weights = embedding.weight weights.grad = torch.zeros((V, D), dtype=dtype, device=device) @@ -57,48 +58,68 @@ def test_embedding(B=2, M=4096, V=150000, D=4096, transpose=False, bench=False): grad = weights.grad grad.zero_() y = triton_embedding_forward(input_ids, weights.data_ptr(), D, dtype) - output_check(y_ref, y, name='y') + output_check(y_ref, y, name="y") triton_embedding_backward(dy, input_ids, grad.data_ptr(), grad.dtype) - output_check(grad_ref, grad.to(dtype), name='grad') + output_check(grad_ref, grad.to(dtype), name="grad") grad.zero_() 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') + output_check(y_ref, y, name="y") + output_check(grad_ref, grad.to(dtype), name="grad") if bench: ref_bytes = B * M * D * 4 - ref_time = benchmark_func(embedding.forward, input_ids, - ref_bytes=ref_bytes) - benchmark_func(embedding_lookup, input_ids, weights, - ref_time=ref_time, ref_bytes=ref_bytes) - - ref_time = benchmark_func(y_ref.backward, dy, retain_graph=True, - ref_bytes=ref_bytes) - benchmark_func(triton_atomic_embedding_backward, dy, input_ids, - grad.data_ptr(), grad.dtype, - ref_time=ref_time, ref_bytes=ref_bytes) - benchmark_func(triton_sync_embedding_backward, dy, input_ids, - grad.data_ptr(), grad.dtype, - ref_time=ref_time, ref_bytes=ref_bytes) - benchmark_func(triton_embedding_backward, dy, input_ids, - grad.data_ptr(), grad.dtype, - ref_time=ref_time, ref_bytes=ref_bytes) - benchmark_func(y.backward, dy, retain_graph=True, - ref_time=ref_time, ref_bytes=ref_bytes) - - -def test_fused_embedding(B=2, M=4096, V=150000, D=4096, use_main_grad=True, - transpose=False, bench=False): + ref_time = benchmark_func(embedding.forward, input_ids, ref_bytes=ref_bytes) + benchmark_func( + embedding_lookup, input_ids, weights, ref_time=ref_time, ref_bytes=ref_bytes + ) + + ref_time = benchmark_func( + y_ref.backward, dy, retain_graph=True, ref_bytes=ref_bytes + ) + benchmark_func( + triton_atomic_embedding_backward, + dy, + input_ids, + grad.data_ptr(), + grad.dtype, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + benchmark_func( + triton_sync_embedding_backward, + dy, + input_ids, + grad.data_ptr(), + grad.dtype, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + benchmark_func( + triton_embedding_backward, + dy, + input_ids, + grad.data_ptr(), + grad.dtype, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) + benchmark_func( + y.backward, dy, retain_graph=True, ref_time=ref_time, ref_bytes=ref_bytes + ) + + +def test_fused_embedding( + B=2, M=4096, V=150000, D=4096, use_main_grad=True, transpose=False, bench=False +): dtype = torch.bfloat16 - device = 'cuda:0' - grad_name = 'main_grad' if use_main_grad else 'grad' + device = "cuda:0" + grad_name = "main_grad" if use_main_grad else "grad" embedding = torch.nn.Embedding(V, D, dtype=dtype, device=device) - input_ids = torch.randint(0, V // 15, (B, M), dtype=torch.int32, - device=device) + input_ids = torch.randint(0, V // 15, (B, M), dtype=torch.int32, device=device) weights = embedding.weight weights.grad = torch.zeros((V, D), dtype=dtype, device=device) if use_main_grad: @@ -117,39 +138,40 @@ def test_fused_embedding(B=2, M=4096, V=150000, D=4096, use_main_grad=True, grad.zero_() y = triton_embedding_forward(input_ids, weights.data_ptr(), D, dtype) - output_check(y_ref, y, name='y') + output_check(y_ref, y, name="y") triton_embedding_backward(dy, input_ids, grad.data_ptr(), grad.dtype) - output_check(grad_ref, grad.to(dtype), name='grad') + output_check(grad_ref, grad.to(dtype), name="grad") grad.zero_() - y = fused_accumulation_embedding_lookup(input_ids, weights, - grad_name=grad_name) + 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, grad.to(dtype), name='grad') + output_check(y_ref, y, name="y") + output_check(grad_ref, grad.to(dtype), name="grad") if bench: ref_bytes = B * M * D * 4 ref_time = benchmark_func(embedding.forward, input_ids) - benchmark_func(fused_accumulation_embedding_lookup, input_ids, weights, - grad_name=grad_name, - ref_time=ref_time, ref_bytes=ref_bytes) + benchmark_func( + fused_accumulation_embedding_lookup, + input_ids, + weights, + grad_name=grad_name, + ref_time=ref_time, + ref_bytes=ref_bytes, + ) ref_time = benchmark_func(y_ref.backward, dy, retain_graph=True) - benchmark_func(y.backward, dy, retain_graph=True, - ref_time=ref_time, ref_bytes=ref_bytes) + benchmark_func( + y.backward, dy, retain_graph=True, ref_time=ref_time, ref_bytes=ref_bytes + ) -if __name__ == '__main__': +if __name__ == "__main__": test_scan(M=8192, bench=False) test_embedding(B=1, M=8192, V=150000, D=8192, transpose=False, bench=False) test_embedding(B=2, M=4096, V=150000, D=4096, transpose=True, bench=False) - test_fused_embedding(B=1, M=8192, V=150000, D=8192, transpose=False, - bench=False) - test_fused_embedding(B=1, M=4096, V=150000, D=8192, transpose=False, - bench=False) - test_fused_embedding(B=2, M=4096, V=150000, D=8192, transpose=True, - bench=False) - test_fused_embedding(B=0, M=4096, V=150000, D=8192, transpose=True, - bench=False) + test_fused_embedding(B=1, M=8192, V=150000, D=8192, transpose=False, bench=False) + test_fused_embedding(B=1, M=4096, V=150000, D=8192, transpose=False, bench=False) + test_fused_embedding(B=2, M=4096, V=150000, D=8192, transpose=True, bench=False) + test_fused_embedding(B=0, M=4096, V=150000, D=8192, transpose=True, bench=False) diff --git a/tests/test_fp32_gemm.py b/tests/test_fp32_gemm.py index 20f9714..6fea714 100644 --- a/tests/test_fp32_gemm.py +++ b/tests/test_fp32_gemm.py @@ -6,19 +6,22 @@ import torch 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.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.tools.benchmark import benchmark_func 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) + return torch.nn.functional.linear(x.to(torch.float64), w.to(torch.float64)).to( + torch.float32 + ) def torch_fp32_matmul(x, w): @@ -35,7 +38,7 @@ def torch_fp32_matmul_update(dy, x): def test_fp32_matmul(M=2048, N=256, K=8192, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" x = torch.randn(M, K, dtype=dtype, device=device, requires_grad=True) w = torch.randn(N, K, dtype=dtype, device=device, requires_grad=True) @@ -50,16 +53,16 @@ def test_fp32_matmul(M=2048, N=256, K=8192, bench=False): dx = triton_fp32_gemm_for_backward(dy, w) dw = triton_fp32_gemm_for_update(dy, x) - 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) + 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) y = triton_split_fp32_gemm(x, w) 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-2) + 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-2) x.grad = None w.grad = None @@ -67,50 +70,88 @@ def test_fp32_matmul(M=2048, N=256, K=8192, bench=False): y.backward(gradient=dy) dx = x.grad dw = w.grad - 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) + 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) if bench: ref_bytes = M * K * 6 + N * K * 6 + M * N * 4 ref_flops = 2 * M * N * K - ref_time = benchmark_func(torch_fp32_matmul, x, w, - ref_bytes=ref_bytes, - ref_flops=ref_flops) - benchmark_func(triton_fp32_gemm, x, w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, ref_time=ref_time) - benchmark_func(triton_split_fp32_gemm, x, w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, ref_time=ref_time) + ref_time = benchmark_func( + torch_fp32_matmul, x, w, ref_bytes=ref_bytes, ref_flops=ref_flops + ) + benchmark_func( + triton_fp32_gemm, + x, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + benchmark_func( + triton_split_fp32_gemm, + x, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) ref_bytes = M * K * 10 + N * K * 4 + M * N * 4 - ref_time = benchmark_func(torch_fp32_matmul_backward, dy, w.float(), - ref_bytes=ref_bytes, - ref_flops=ref_flops) - benchmark_func(triton_fp32_gemm_for_backward, dy, w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, ref_time=ref_time) - benchmark_func(triton_split_fp32_gemm_for_backward, dy, w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, ref_time=ref_time) + ref_time = benchmark_func( + torch_fp32_matmul_backward, + dy, + w.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ) + benchmark_func( + triton_fp32_gemm_for_backward, + dy, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + benchmark_func( + triton_split_fp32_gemm_for_backward, + dy, + w, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) ref_bytes = M * K * 4 + N * K * 12 + M * N * 4 - ref_time = benchmark_func(torch_fp32_matmul_update, dy, x.float(), - ref_bytes=ref_bytes, - ref_flops=ref_flops) - benchmark_func(triton_fp32_gemm_for_update, dy, x, - ref_bytes=ref_bytes, - ref_flops=ref_flops, ref_time=ref_time) - benchmark_func(triton_split_fp32_gemm_for_update, dy, x, - ref_bytes=ref_bytes, - ref_flops=ref_flops, ref_time=ref_time) + ref_time = benchmark_func( + torch_fp32_matmul_update, + dy, + x.float(), + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ) + benchmark_func( + triton_fp32_gemm_for_update, + dy, + x, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) + benchmark_func( + triton_split_fp32_gemm_for_update, + dy, + x, + ref_bytes=ref_bytes, + ref_flops=ref_flops, + ref_time=ref_time, + ) def test_BMK_fp32_matmul(B=2, M=2048, N=256, K=8192, bench=False): # M, N, K = 4096, 256, 8192 dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 x = torch.randn(B, M, K, dtype=dtype, device=device, requires_grad=True) @@ -128,21 +169,32 @@ def test_BMK_fp32_matmul(B=2, M=2048, N=256, K=8192, bench=False): 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=1e-1, rtol=2e-2) - output_check(dw_ref, dw, name='update', atol=1e-1, rtol=2e-2) + output_check(y_ref, y, name="forward", atol=5e-3, rtol=2e-3) + output_check(dx_ref, dx, name="backward", atol=1e-1, rtol=2e-2) + output_check(dw_ref, dw, name="update", atol=1e-1, rtol=2e-2) 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_flops=2 * M * N * K) - benchmark_func(fp32_gemm, x, w, 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) - - -if __name__ == '__main__': + 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_flops=2 * M * N * K, + ) + benchmark_func( + fp32_gemm, + x, + w, + 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, + ) + + +if __name__ == "__main__": test_fp32_matmul(M=4096, N=256, K=8192, bench=False) test_fp32_matmul(M=16384, N=256, K=2048, bench=False) test_fp32_matmul(M=128, N=16, K=128, bench=False) diff --git a/tests/test_gate.py b/tests/test_gate.py index 34fcb5a..3cd8baf 100644 --- a/tests/test_gate.py +++ b/tests/test_gate.py @@ -8,13 +8,16 @@ from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check -from linghe.utils.gate import (triton_group_rms_norm_gate_forward, - triton_group_rms_norm_gate_backward) +from linghe.utils.gate import ( + triton_group_rms_norm_gate_forward, + triton_group_rms_norm_gate_backward, +) # @torch.compile -def torch_group_rms_norm_gate_forward(x, gate, weight, eps=1e-6, group_size=4, - transpose=True): +def torch_group_rms_norm_gate_forward( + x, gate, weight, eps=1e-6, group_size=4, transpose=True +): dtype = x.dtype x = x.float() gate = gate.float() @@ -28,11 +31,11 @@ def torch_group_rms_norm_gate_forward(x, gate, weight, eps=1e-6, group_size=4, 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) + 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) + o = F.rms_norm(attn_output[:, :, i], [d], weight=weight, eps=eps) outputs.append(o) outputs = torch.stack(outputs, 2).view(bs, length, dim) if transpose: @@ -42,99 +45,144 @@ def torch_group_rms_norm_gate_forward(x, gate, weight, eps=1e-6, group_size=4, return outputs -def torch_group_rms_norm_gate_backward(grad_output, x, gate, weight, eps=1e-6, - group_size=4, - transpose=True): +def torch_group_rms_norm_gate_backward( + grad_output, x, gate, weight, eps=1e-6, group_size=4, transpose=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, - transpose=transpose) + 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.to(dtype), gate.grad.to(dtype), weight.grad.to(dtype) -def test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, - transpose=True, share=False, coef=1.0, - grad_coef=1.0, - bench=False): +def test_group_rms_norm_gate( + bs=1, + length=4096, + dim=4096, + group_size=4, + transpose=True, + share=False, + coef=1.0, + grad_coef=1.0, + bench=False, +): dtype = torch.bfloat16 - device = 'cuda:0' - x = torch.randn(bs, length, 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) + device = "cuda:0" + x = torch.randn(bs, length, 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 transpose: - gate = (torch.randn(length, bs, dim, dtype=dtype, - device=device) * coef).requires_grad_() - grad_output = torch.randn(length, bs, dim, dtype=dtype, - device=device) * grad_coef + gate = ( + torch.randn(length, bs, dim, dtype=dtype, device=device) * coef + ).requires_grad_() + grad_output = ( + torch.randn(length, bs, dim, dtype=dtype, device=device) * grad_coef + ) else: - gate = (torch.randn(bs, length, dim, dtype=dtype, - device=device) * coef).requires_grad_() - grad_output = torch.randn(bs, length, dim, dtype=dtype, - device=device) * grad_coef + gate = ( + torch.randn(bs, length, dim, dtype=dtype, device=device) * coef + ).requires_grad_() + grad_output = ( + torch.randn(bs, length, dim, dtype=dtype, device=device) * grad_coef + ) - 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, name='group_norm_gate.y') + 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, name="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, 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') + 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, 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") if bench: - benchmark_func(torch_group_rms_norm_gate_forward, x, gate, weight, - group_size=group_size, transpose=transpose, - ref_bytes=bs * length * dim * 6) + benchmark_func( + torch_group_rms_norm_gate_forward, + x, + gate, + weight, + group_size=group_size, + transpose=transpose, + ref_bytes=bs * length * dim * 6, + ) - benchmark_func(triton_group_rms_norm_gate_forward, x, gate, weight, - group_size=group_size, transpose=transpose, - ref_bytes=bs * length * dim * 6) + benchmark_func( + triton_group_rms_norm_gate_forward, + x, + gate, + weight, + group_size=group_size, + transpose=transpose, + ref_bytes=bs * length * dim * 6, + ) - benchmark_func(triton_group_rms_norm_gate_backward, grad_output, x, - gate, - weight, group_size=group_size, transpose=transpose, - ref_bytes=bs * length * dim * 10) + benchmark_func( + triton_group_rms_norm_gate_backward, + grad_output, + x, + gate, + weight, + group_size=group_size, + transpose=transpose, + ref_bytes=bs * length * dim * 10, + ) -if __name__ == '__main__': - 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=2, length=4096, dim=2048, group_size=4, - transpose=False, share=True, - bench=False) - test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, - bench=False) - test_group_rms_norm_gate(bs=2, length=4096, dim=1536, group_size=4, - transpose=True, - bench=False) - test_group_rms_norm_gate(bs=2, length=4096, dim=1536, group_size=4, - transpose=False, - coef=10000.0, - grad_coef=10000.0, - bench=False) - test_group_rms_norm_gate(bs=2, length=4096, dim=1536, group_size=4, - transpose=False, - coef=0.0, - grad_coef=0.0, - bench=False) +if __name__ == "__main__": + 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=2, + length=4096, + dim=2048, + group_size=4, + transpose=False, + share=True, + bench=False, + ) + test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, bench=False) + test_group_rms_norm_gate( + bs=2, length=4096, dim=1536, group_size=4, transpose=True, bench=False + ) + test_group_rms_norm_gate( + bs=2, + length=4096, + dim=1536, + group_size=4, + transpose=False, + coef=10000.0, + grad_coef=10000.0, + bench=False, + ) + test_group_rms_norm_gate( + bs=2, + length=4096, + dim=1536, + group_size=4, + transpose=False, + coef=0.0, + grad_coef=0.0, + bench=False, + ) diff --git a/tests/test_gather.py b/tests/test_gather.py index 035b3ae..c47aeef 100644 --- a/tests/test_gather.py +++ b/tests/test_gather.py @@ -8,20 +8,23 @@ from linghe.quant.block import triton_batch_blockwise_quant from linghe.tools.benchmark import benchmark_func 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) -from linghe.utils.gather import (triton_make_row_id_map, - triton_make_row_id_map_and_index, - 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, - triton_batch_block_pad_permute_with_indices, - ) +from linghe.tools.util import ( + torch_batch_smooth_quant, + torch_blockwise_quant, + torch_make_indices, + torch_smooth_quant, +) +from linghe.utils.gather import ( + triton_make_row_id_map, + triton_make_row_id_map_and_index, + 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, + triton_batch_block_pad_permute_with_indices, +) def torch_index_select(y, indices): @@ -32,8 +35,7 @@ def torch_index_select(y, indices): 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) + 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) for i in range(E): @@ -64,28 +66,30 @@ def torch_scatter(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): +def torch_smooth_permute_with_indices( + grad_data, + grad_scale, + indices, + 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]) + B = grad_data.shape[1] // (1 if grad_scale.ndim == 1 else grad_scale.shape[1]) q_refs = [] scale_refs = [] s = 0 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) + 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] + 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] scale = y_smooth.abs().amax(1) / 448 @@ -102,12 +106,15 @@ def torch_smooth_permute_with_indices(grad_data, grad_scale, indices, # 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): +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, +): M, DIM = x_q.shape q_refs = [] scale_refs = [] @@ -115,24 +122,23 @@ def torch_batch_transpose_smooth_permute_with_indices(x_q, x_scale, 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_q.device) + 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]] + 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_q, y_scale, y_max = 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 @@ -141,11 +147,9 @@ def torch_batch_transpose_smooth_permute_with_indices(x_q, x_scale, return q_ref, scale_ref -def torch_batch_block_pad_permute_with_indices(x, - indices, - probs, - token_count_per_expert_list, - round_scale=True): +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 @@ -166,7 +170,7 @@ def torch_batch_block_pad_permute_with_indices(x, c = token_count_per_expert_list[i] if c == 0: continue - index = indices[s:s + c] + index = indices[s : s + c] assert len(index) == c y = x[index] y = y.float() @@ -176,9 +180,9 @@ def torch_batch_block_pad_permute_with_indices(x, 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) + 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)) @@ -195,11 +199,12 @@ def torch_batch_block_pad_permute_with_indices(x, def test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): 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) @@ -211,47 +216,66 @@ def test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): 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' +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) + 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) + logits, topk=topk, bias=0.0 + ) - tokens = torch.randn((indices.shape[0], N), dtype=torch.bfloat16, - device=device) + 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) + 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') + 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' + 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" dtype = torch.bfloat16 x = torch.randn(M, N, dtype=dtype, device=device) ** 3 x_q = x.to(torch.float8_e4m3fn) @@ -260,23 +284,22 @@ 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_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') + 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) @@ -284,276 +307,391 @@ def test_triton_permute_with_mask_map(M=4096, N=4096, n_experts=256, topk=8, 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') + [(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, - tokens_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' + 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, + tokens_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" 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 + 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) + 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=dtype, device=device).to( - torch.float8_e4m3fn) + grad_data = torch.randn((M, N), dtype=dtype, 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, y_q, name='data', rtol=0.125) - output_check(scale_ref, y_scale, name='scale') + 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, y_q, name="data", rtol=0.125) + output_check(scale_ref, y_scale, name="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(), name='smoothed.data', - rtol=0.125) - 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) + 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(), name="smoothed.data", rtol=0.125) + 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(), name='smoothed.data', - rtol=0.125) - output_check(scale_ref.float(), permuted_scale.float(), 'smoothed.scale') + 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(), name="smoothed.data", rtol=0.125) + 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' + 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" if True: - logits = torch.randn((M, n_experts), dtype=torch.float32, - device=device) ** 3 + 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) + 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 + 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 + ) 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_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(), name='smoothed.data', rtol=0.125) - 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') + 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(), name="smoothed.data", rtol=0.125) + 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) - - -def test_batch_block_pad_permute_with_indices(M=16384, N=2048, n_experts=32, - topk=2, bench=False): - device = 'cuda:0' - logits = torch.randn((M, n_experts), dtype=torch.float32, - device=device) ** 3 + 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, + ) + + +def test_batch_block_pad_permute_with_indices( + M=16384, N=2048, n_experts=32, topk=2, bench=False +): + 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) + 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) + 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_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, - indices, - probs, + token_count_per_expert, + pad_indices, 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') + 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") if bench: - benchmark_func(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_func(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) + benchmark_func( + 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_func( + 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_func(triton_batch_blockwise_quant, xs, token_count_per_expert, - token_count_per_expert_list, - round_scale=True, - ref_bytes=num_out_tokens * N * 4) - + benchmark_func( + triton_batch_blockwise_quant, + xs, + token_count_per_expert, + token_count_per_expert_list, + round_scale=True, + ref_bytes=num_out_tokens * N * 4, + ) -if __name__ == '__main__': +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) - - test_batch_block_pad_permute_with_indices(M=8192 * 2, N=2048, n_experts=32, - topk=2, bench=False) - test_batch_block_pad_permute_with_indices(M=0, N=2048, n_experts=32, topk=2, - bench=False) - test_batch_block_pad_permute_with_indices(M=8192, N=1536, n_experts=32, - topk=2, 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 + ) + + test_batch_block_pad_permute_with_indices( + M=8192 * 2, N=2048, n_experts=32, topk=2, bench=False + ) + test_batch_block_pad_permute_with_indices( + M=0, N=2048, n_experts=32, topk=2, bench=False + ) + test_batch_block_pad_permute_with_indices( + M=8192, N=1536, n_experts=32, topk=2, bench=False + ) diff --git a/tests/test_group_quant.py b/tests/test_group_quant.py index b7c4c26..d9d96dd 100644 --- a/tests/test_group_quant.py +++ b/tests/test_group_quant.py @@ -12,19 +12,20 @@ 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 + 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, xq, name='data') - output_check(x_scale_ref, x_scale, name='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) + benchmark_func( + triton_group_quant, x, group_size=B, n_repeat=n_repeat, ref_bytes=M * N * 3 + ) -if __name__ == '__main__': +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) diff --git a/tests/test_hadamard_quant.py b/tests/test_hadamard_quant.py index ae9263b..12880e9 100644 --- a/tests/test_hadamard_quant.py +++ b/tests/test_hadamard_quant.py @@ -8,10 +8,11 @@ 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, - ) +from linghe.tools.util import ( + make_hadamard_matrix, + torch_hadamard_transform, + torch_row_quant, +) # apply hadamard transformation and quantization for x @@ -19,9 +20,9 @@ def torch_hadamard_quant(x, hm, round_scale=False): dtype = x.dtype x = x.float() hm = hm.float() - xh = torch_hadamard_transform(x, hm, side='right') + 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') + xht = torch_hadamard_transform(x.t().contiguous(), hm, side="right") qt, st = torch_row_quant(xht, round_scale=round_scale) return xh.to(dtype), xht.to(dtype), q, s, qt, st @@ -29,7 +30,7 @@ def torch_hadamard_quant(x, hm, round_scale=False): def test_hadamard_quant(M=8192, N=1024, K=2048, B=64, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + 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) @@ -42,39 +43,38 @@ def test_hadamard_quant(M=8192, N=1024, K=2048, B=64, bench=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) + 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() - 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) + 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', 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') + 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', 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') + 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', 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') + 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' + 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) @@ -83,17 +83,17 @@ def test_hadamard_quant_linear(M=8192, N=1024, K=2048, B=64): y_ref = x @ w.t() y = linear(x) - output_check(y_ref, y, name='y', rtol=-1) + output_check(y_ref, y, name="y", rtol=-1) dx_ref = dy @ w dw_ref = dy.t() @ x y.backward(dy) dw = linear.weight.grad dx = x.grad - output_check(dx_ref, dx, name='dx', rtol=-1) - output_check(dw_ref, dw, name='dw', rtol=-1) + 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) diff --git a/tests/test_la.py b/tests/test_la.py index 6661e32..32f681f 100644 --- a/tests/test_la.py +++ b/tests/test_la.py @@ -7,9 +7,11 @@ import torch -from linghe.attn.la import (triton_lightning_attention_forward, - triton_lightning_attention_backward, - triton_fused_lightning_attention_backward) +from linghe.attn.la import ( + triton_lightning_attention_forward, + triton_lightning_attention_backward, + triton_fused_lightning_attention_backward, +) from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check @@ -46,24 +48,24 @@ def torch_la(q, k, v, s, decay_scales): 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() + 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)) + 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): +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') +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 = [] @@ -84,21 +86,19 @@ def make_varlen_input(qo_heads=16, kv_heads=16, dim=128, qls=[1024, 1024], return q, k, v -def test_la(bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, - bench=False): - device = torch.device('cuda:0') +def test_la( + bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, bench=False +): + device = torch.device("cuda:0") dtype = torch.bfloat16 - q = torch.randn(bs, length, qo_heads, dim, dtype=dtype, - device=device) ** 3 * 0.1 + 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 = 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 = 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) @@ -106,8 +106,8 @@ def test_la(bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, 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)) + -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) @@ -120,32 +120,53 @@ def test_la(bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, 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) + 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) + 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) + 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) if bench: ref_bytes = bs * length * qo_heads * dim * 8 + bs * qo_heads * dim * dim * 8 - benchmark_func(triton_lightning_attention_forward, q, k, v, - decay_scales, ref_bytes=ref_bytes) - benchmark_func(triton_lightning_attention_backward, g, q, k, v, - decay_scales, ref_bytes=ref_bytes) - benchmark_func(triton_fused_lightning_attention_backward, g, q, k, v, - state, decay_scales, ref_bytes=ref_bytes) + benchmark_func( + triton_lightning_attention_forward, + q, + k, + v, + decay_scales, + ref_bytes=ref_bytes, + ) + benchmark_func( + triton_lightning_attention_backward, + g, + q, + k, + v, + decay_scales, + ref_bytes=ref_bytes, + ) + benchmark_func( + triton_fused_lightning_attention_backward, + g, + q, + k, + v, + state, + decay_scales, + ref_bytes=ref_bytes, + ) # def test_varlen_la(qls=[1024,1024], qo_heads=16, kv_heads=16, dim=128, digest=False, bench=False): @@ -195,6 +216,7 @@ def test_la(bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, # benchmark_func(triton_lightning_attention_forward, q, k, v, decay_scales, cu_seqlens, padded_cu_seqlens, max_q_length, ref_bytes=ref_bytes) -if __name__ == '__main__': - test_la(bs=1, length=8192, qo_heads=64, kv_heads=64, dim=128, digest=False, - bench=False) +if __name__ == "__main__": + test_la( + bs=1, length=8192, qo_heads=64, kv_heads=64, dim=128, digest=False, bench=False + ) diff --git a/tests/test_loss.py b/tests/test_loss.py index a95fb1b..d3a9543 100644 --- a/tests/test_loss.py +++ b/tests/test_loss.py @@ -10,10 +10,12 @@ from linghe.facade.loss import moe_z_loss, softmax_cross_entropy from linghe.tools.benchmark import benchmark_func 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) +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, grad, ignore_index=-100): @@ -21,8 +23,9 @@ def torch_cross_entropy(logits, targets, grad, ignore_index=-100): losses = torch.nn.functional.cross_entropy( float_logits.view(-1, logits.size()[-1]), targets.view(-1), - reduction='none', - ignore_index=ignore_index) + reduction="none", + ignore_index=ignore_index, + ) loss = (losses * grad).sum() loss.backward() return losses.to(torch.float32), logits.grad @@ -30,19 +33,24 @@ def torch_cross_entropy(logits, targets, grad, ignore_index=-100): 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 = torch.mean(torch.square(torch.logsumexp(float_logits, dim=-1))) * coef loss.backward() return loss, logits.grad -def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, grad_coef=1.0, - ignore_index=None, fill=False, - inplace=False, bench=False): - device = 'cuda:0' +def test_triton_softmax_cross_entropy( + M=4096, + N=157184, + coef=1.0, + grad_coef=1.0, + ignore_index=None, + fill=False, + inplace=False, + bench=False, +): + device = "cuda:0" dtype = torch.bfloat16 - logits = torch.randn((M, N), dtype=dtype, device=device, - requires_grad=False) + logits = torch.randn((M, N), dtype=dtype, device=device, requires_grad=False) select = True if select: @@ -67,103 +75,165 @@ def test_triton_softmax_cross_entropy(M=4096, N=157184, coef=1.0, grad_coef=1.0, ignore_index = -100 if ignore_index is None else ignore_index logits = (logits * coef).detach().clone().requires_grad_() - 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) + 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, - 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) + 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 = 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) + output_check(loss_ref, loss, name="loss", atol=1e-4, rtol=1e-5) + output_check(grad_ref, grad, name="grad", digest=10) if bench: - benchmark_func(torch_cross_entropy, logits.requires_grad_(), targets, - output_grad, - ref_bytes=M * N * 2) - benchmark_func(triton_softmax_cross_entropy_forward, logits, targets, - ignore_index=ignore_index, - ref_bytes=M * N * 2) - benchmark_func(triton_softmax_cross_entropy_backward, - logits.detach().clone(), targets, - sum_exp, max_logit, output_grad, - ignore_index=ignore_index, ref_bytes=M * N * 4) + benchmark_func( + torch_cross_entropy, + logits.requires_grad_(), + targets, + output_grad, + ref_bytes=M * N * 2, + ) + benchmark_func( + triton_softmax_cross_entropy_forward, + logits, + targets, + ignore_index=ignore_index, + ref_bytes=M * N * 2, + ) + benchmark_func( + triton_softmax_cross_entropy_backward, + logits.detach().clone(), + targets, + sum_exp, + max_logit, + output_grad, + ignore_index=ignore_index, + ref_bytes=M * N * 4, + ) def test_z_loss(L=4096, B=2, N=256, coef=0.001, bench=False): - device = 'cuda:0' - logits = torch.randn((L, B, N), dtype=torch.float32, device=device, - requires_grad=False) + 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') + 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') + output_check(loss_ref, loss, name="loss") + output_check(grad_ref.float(), grad.float(), name="grad") if bench: - benchmark_func(torch_z_loss, logits, coef=coef, - ref_bytes=L * B * N * 4) - benchmark_func(triton_moe_z_loss_forward, logits, coef=coef, - ref_bytes=L * B * N * 4) - benchmark_func(triton_moe_z_loss_backward, input_grad, logits, - coef=coef, ref_bytes=L * B * N * 8) - - -if __name__ == '__main__': - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, - inplace=True, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, - grad_coef=1e-6, inplace=True, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=10000.0, - grad_coef=100.0, fill=True, inplace=True, - bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, - fill=True, ignore_index=-100, - inplace=True, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1.0, grad_coef=1.0, - fill=True, ignore_index=0, inplace=True, - bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184 - 16, coef=10000.0, - grad_coef=100.0, fill=True, inplace=True, - bench=False) - test_triton_softmax_cross_entropy(M=8192, N=175175, coef=1.0, grad_coef=1.0, - inplace=True, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=0.0, grad_coef=0.0, - inplace=True, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=0.0, - grad_coef=100.0, inplace=True, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=1000.0, - grad_coef=0.0, inplace=True, bench=False) - test_triton_softmax_cross_entropy(M=8192, N=157184, coef=100.0, - grad_coef=100.0, fill=True, inplace=True, - bench=False) - test_triton_softmax_cross_entropy(M=4096, N=157184, coef=0.1, grad_coef=1.0, - inplace=True, bench=False) + benchmark_func(torch_z_loss, logits, coef=coef, ref_bytes=L * B * N * 4) + benchmark_func( + triton_moe_z_loss_forward, logits, coef=coef, ref_bytes=L * B * N * 4 + ) + benchmark_func( + triton_moe_z_loss_backward, + input_grad, + logits, + coef=coef, + ref_bytes=L * B * N * 8, + ) + + +if __name__ == "__main__": + test_triton_softmax_cross_entropy( + M=8192, N=157184, coef=1.0, grad_coef=1.0, inplace=True, bench=False + ) + test_triton_softmax_cross_entropy( + M=8192, N=157184, coef=1.0, grad_coef=1e-6, inplace=True, bench=False + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184, + coef=10000.0, + grad_coef=100.0, + fill=True, + inplace=True, + bench=False, + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184, + coef=1.0, + grad_coef=1.0, + fill=True, + ignore_index=-100, + inplace=True, + bench=False, + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184, + coef=1.0, + grad_coef=1.0, + fill=True, + ignore_index=0, + inplace=True, + bench=False, + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184 - 16, + coef=10000.0, + grad_coef=100.0, + fill=True, + inplace=True, + bench=False, + ) + test_triton_softmax_cross_entropy( + M=8192, N=175175, coef=1.0, grad_coef=1.0, inplace=True, bench=False + ) + test_triton_softmax_cross_entropy( + M=8192, N=157184, coef=0.0, grad_coef=0.0, inplace=True, bench=False + ) + test_triton_softmax_cross_entropy( + M=8192, N=157184, coef=0.0, grad_coef=100.0, inplace=True, bench=False + ) + test_triton_softmax_cross_entropy( + M=8192, N=157184, coef=1000.0, grad_coef=0.0, inplace=True, bench=False + ) + test_triton_softmax_cross_entropy( + M=8192, + N=157184, + coef=100.0, + grad_coef=100.0, + fill=True, + inplace=True, + bench=False, + ) + test_triton_softmax_cross_entropy( + M=4096, N=157184, coef=0.1, grad_coef=1.0, inplace=True, bench=False + ) test_z_loss(L=4096, B=2, N=256, coef=1e-6, bench=False) diff --git a/tests/test_mla.py b/tests/test_mla.py index 74c8a74..e234b4a 100644 --- a/tests/test_mla.py +++ b/tests/test_mla.py @@ -7,11 +7,13 @@ 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.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.benchmark import benchmark_func from linghe.tools.check import output_check @@ -53,13 +55,15 @@ def torch_attn(q, k, v, causal=True, mask=None, clip_value=None, hp=False): 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() + 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): +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() @@ -87,7 +91,8 @@ def torch_varlen_attn(qs, ks, vs, cu_seqlens, padded_cu_seqlens=None, 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]) + 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])) @@ -120,14 +125,13 @@ def head_wise_quant(x): def test_softmax(M=128, N=128): - 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') + 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') + output_check(grad_ref, grad, atol=10, name="grad") def test_dot_sum(M=128, N=128, D=128): @@ -136,25 +140,32 @@ def test_dot_sum(M=128, N=128, D=128): 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') - - -def test_mla(B=2, L=4096, H=16, causal=True, hpc=False, safe=True, coef=1.0, - clip_value=None, bench=False): + output_check(ds_ref, ds, atol=10, name="dot_sum") + + +def test_mla( + B=2, + L=4096, + H=16, + causal=True, + hpc=False, + safe=True, + coef=1.0, + clip_value=None, + bench=False, +): dtype = torch.bfloat16 - device = 'cuda:0' - q = (torch.randn((B, L, H, 192), device=device, - dtype=dtype) * coef).requires_grad_() + 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) + 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, 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 @@ -164,50 +175,85 @@ def test_mla(B=2, L=4096, H=16, causal=True, hpc=False, safe=True, coef=1.0, 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, 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) + 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') + 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") if bench: ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) - benchmark_func(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_func(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) - - -def test_varlen_mla(LS=[2048, 4096], H=16, causal=True, hpc=False, safe=True, - coef=1.0, clip_value=None, pad=False, bench=False): + benchmark_func( + 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_func( + 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, + ) + + +def test_varlen_mla( + LS=[2048, 4096], + H=16, + causal=True, + hpc=False, + safe=True, + coef=1.0, + clip_value=None, + pad=False, + bench=False, +): dtype = torch.bfloat16 - device = 'cuda:0' + 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) + 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_() + 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_() @@ -216,9 +262,9 @@ def test_varlen_mla(LS=[2048, 4096], H=16, causal=True, hpc=False, safe=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, 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 @@ -228,99 +274,144 @@ def test_varlen_mla(LS=[2048, 4096], H=16, causal=True, hpc=False, safe=True, 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') + 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') + 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) + 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=clip_value) + 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=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') + 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') + 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') + 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") if bench: + ref_flops = sum([L * L * H * (192 + 128) * (1 if causal else 2) for L in LS]) + benchmark_func( + 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) * (1 if causal else 2) for L in LS]) - benchmark_func(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_func(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) - - -def test_fp8_mla(B=2, L=4096, H=16, causal=True, hpc=False, quant_value=False, - bench=False): + [L * L * H * (192 + 128 * 2 + 192 * 2) * (1 if causal else 2) for L in LS] + ) + benchmark_func( + 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, + ) + + +def test_fp8_mla( + B=2, L=4096, H=16, causal=True, hpc=False, quant_value=False, bench=False +): 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) + 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_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') + 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") if bench: ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) - benchmark_func(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) + benchmark_func( + 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, + ) if __name__ == "__main__": @@ -328,25 +419,180 @@ def test_fp8_mla(B=2, L=4096, H=16, causal=True, hpc=False, quant_value=False, test_dot_sum(M=128, N=128, D=128) - test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=False, coef=1.0, - clip_value=500, bench=False) - test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=500.0, bench=False) - test_mla(B=1, L=8192, H=64, causal=True, hpc=True, safe=False, coef=1.0, clip_value=None, bench=False) - test_mla(B=1, L=8192, H=64, causal=True, hpc=False, safe=True, coef=100.0, clip_value=None, bench=False) - test_mla(B=1, L=4096, H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - test_mla(B=1, L=4096, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - test_mla(B=1, L=8192, H=64, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - test_mla(B=1, L=8192, H=1, causal=False, hpc=False, safe=False, coef=1.0, clip_value=None, bench=False) - - test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=False, coef=1.0, clip_value=None, pad=False, bench=False) - test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=100.0, pad=False, bench=False) - test_varlen_mla(LS=[8192], H=64, causal=True, hpc=False, safe=True, coef=1.0, clip_value=None, pad=True, bench=False) - - test_varlen_mla(LS=[4096,4096], H=64, causal=True, hpc=False, safe=True, coef=1.0, bench=False) - test_varlen_mla(LS=[2048,2048,4096], H=64, causal=True, hpc=True, safe=True, coef=1.0, bench=False) - test_varlen_mla(LS=[127,873,3096], H=64, causal=False, hpc=False, safe=False, coef=1.0, bench=False) - test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, clip_value=100.0, bench=False) - test_varlen_mla(LS=[127,873,3456], H=16, causal=False, hpc=False, safe=True, coef=1.0, pad=True, bench=False) - test_varlen_mla(LS=[127,873,3456], H=1, causal=True, hpc=False, safe=True, coef=1.0, bench=False) - - test_fp8_mla(B=1, L=8192, H=64, causal=True, hpc=False, quant_value=False, bench=False) + test_mla( + B=1, + L=8192, + H=64, + causal=True, + hpc=False, + safe=False, + coef=1.0, + clip_value=500, + bench=False, + ) + test_mla( + B=1, + L=8192, + H=64, + causal=True, + hpc=False, + safe=False, + coef=1.0, + clip_value=500.0, + bench=False, + ) + test_mla( + B=1, + L=8192, + H=64, + causal=True, + hpc=True, + safe=False, + coef=1.0, + clip_value=None, + bench=False, + ) + test_mla( + B=1, + L=8192, + H=64, + causal=True, + hpc=False, + safe=True, + coef=100.0, + clip_value=None, + bench=False, + ) + test_mla( + B=1, + L=4096, + H=64, + causal=True, + hpc=False, + safe=False, + coef=1.0, + clip_value=None, + bench=False, + ) + test_mla( + B=1, + L=4096, + H=64, + causal=False, + hpc=False, + safe=False, + coef=1.0, + clip_value=None, + bench=False, + ) + test_mla( + B=1, + L=8192, + H=64, + causal=False, + hpc=False, + safe=False, + coef=1.0, + clip_value=None, + bench=False, + ) + test_mla( + B=1, + L=8192, + H=1, + causal=False, + hpc=False, + safe=False, + coef=1.0, + clip_value=None, + bench=False, + ) + + test_varlen_mla( + LS=[8192], + H=64, + causal=True, + hpc=False, + safe=False, + coef=1.0, + clip_value=None, + pad=False, + bench=False, + ) + test_varlen_mla( + LS=[8192], + H=64, + causal=True, + hpc=False, + safe=True, + coef=1.0, + clip_value=100.0, + pad=False, + bench=False, + ) + test_varlen_mla( + LS=[8192], + H=64, + causal=True, + hpc=False, + safe=True, + coef=1.0, + clip_value=None, + pad=True, + bench=False, + ) + + test_varlen_mla( + LS=[4096, 4096], H=64, causal=True, hpc=False, safe=True, coef=1.0, bench=False + ) + test_varlen_mla( + LS=[2048, 2048, 4096], + H=64, + causal=True, + hpc=True, + safe=True, + coef=1.0, + bench=False, + ) + test_varlen_mla( + LS=[127, 873, 3096], + H=64, + causal=False, + hpc=False, + safe=False, + coef=1.0, + bench=False, + ) + test_varlen_mla( + LS=[127, 873, 3456], + H=16, + causal=False, + hpc=False, + safe=True, + coef=1.0, + clip_value=100.0, + bench=False, + ) + test_varlen_mla( + LS=[127, 873, 3456], + H=16, + causal=False, + hpc=False, + safe=True, + coef=1.0, + pad=True, + bench=False, + ) + test_varlen_mla( + LS=[127, 873, 3456], + H=1, + causal=True, + hpc=False, + safe=True, + coef=1.0, + bench=False, + ) + + test_fp8_mla( + B=1, L=8192, H=64, causal=True, hpc=False, quant_value=False, bench=False + ) diff --git a/tests/test_mul.py b/tests/test_mul.py index 104d190..0a10342 100644 --- a/tests/test_mul.py +++ b/tests/test_mul.py @@ -9,8 +9,7 @@ from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check -from linghe.utils.mul import triton_dot, triton_batch_scale, \ - triton_inplace_scale +from linghe.utils.mul import triton_dot, triton_batch_scale, triton_inplace_scale def torch_fp16_dot(x, y): @@ -28,7 +27,7 @@ def torch_batch_scale(xs, scale): def test_dot(M=4096, N=4096, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 @@ -40,65 +39,75 @@ def test_dot(M=4096, N=4096, bench=False): sums_ref = torch_fp16_dot(x, q.float().to(dtype)) sums = triton_dot(x, q) - output_check(sums_ref, sums, 'dot', atol=1.0) + 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) + 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) - ref_time = benchmark_func(triton_dot, x, q, n_repeat=n_repeat, - ref_time=ref_time) + ref_time = benchmark_func( + triton_dot, x, q, n_repeat=n_repeat, ref_time=ref_time + ) -def test_inplace_scale(M=2 ** 20, bench=False): - x = torch.randn((M,), device='cuda:0', dtype=torch.float32) +def test_inplace_scale(M=2**20, bench=False): + 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') + output_check(sum_ref, sums, "sum") ref_bytes = M * 8 if bench: - ref_time = benchmark_func(torch_inplace_scale, x, scale, - ref_bytes=ref_bytes) - benchmark_func(triton_inplace_scale, x, scale, - ref_bytes=ref_bytes, ref_time=ref_time) + ref_time = benchmark_func(torch_inplace_scale, x, scale, ref_bytes=ref_bytes) + benchmark_func( + triton_inplace_scale, x, scale, ref_bytes=ref_bytes, ref_time=ref_time + ) def test_batch_scale(M=4096, N=2048, k=128, scale=1.0, bench=False): 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)] - # xs.append(torch.randn(2**32//N, N, + 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) + ] + # xs.append(torch.randn(2**32//N, N, # dtype=dtype, device='cuda:0')) 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') + 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') + 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 if bench: - ref_time = benchmark_func(torch_batch_scale, xs, scale, - ref_bytes=ref_bytes) - benchmark_func(triton_batch_scale, xs, scale, - ref_bytes=ref_bytes, ref_time=ref_time) + ref_time = benchmark_func(torch_batch_scale, xs, scale, ref_bytes=ref_bytes) + benchmark_func( + triton_batch_scale, xs, scale, ref_bytes=ref_bytes, ref_time=ref_time + ) -if __name__ == '__main__': +if __name__ == "__main__": test_dot(M=4096, N=4096, bench=False) - test_inplace_scale(M=2 ** 28 + 1, bench=False) + test_inplace_scale(M=2**28 + 1, bench=False) test_batch_scale(M=2048, N=1024, k=128, scale=2.0, bench=False) test_batch_scale(M=2048, N=1024, k=128, scale=0.0, bench=False) diff --git a/tests/test_norm.py b/tests/test_norm.py index e3359aa..6745e77 100644 --- a/tests/test_norm.py +++ b/tests/test_norm.py @@ -9,13 +9,14 @@ 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) +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): @@ -24,10 +25,7 @@ def torch_rms_forward(x, weight): 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) @@ -41,10 +39,7 @@ def torch_rms_backward(x, weight, dy): 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) @@ -54,24 +49,21 @@ def torch_rms_backward(x, weight, dy): 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 @@ -80,33 +72,26 @@ 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) + 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.t(), rms, yt_q, yt_scale.t() -def torch_rms_gemm_block_quant_forward(x, norm_weight, route_weight, - round_scale=False): +def torch_rms_gemm_block_quant_forward(x, norm_weight, route_weight, round_scale=False): dtype = x.dtype x = x.float() 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 + normalized_shape=N, eps=1e-6, dtype=torch.float32, device=x.device ) with torch.no_grad(): rmsnorm.weight.copy_(norm_weight) @@ -119,8 +104,7 @@ def torch_rms_gemm_block_quant_forward(x, norm_weight, route_weight, 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): +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) @@ -129,7 +113,7 @@ def split_rms_gemm_block_quant_forward(x, norm_weight, route_weight, 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) @@ -137,176 +121,205 @@ def test_rmsnorm(M=4096, N=4096, bench=False): y_ref = torch_rms_forward(x, weight) y, rms = triton_rms_norm_forward(x, weight) - output_check(y_ref, y, 'y') + 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') + 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, name="dx") - output_check(dw_ref, dw.to(dtype), name='dw') + 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') + 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_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) + 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) + 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') + output_check(scale_ref, scale, name="smooth.scale") if calibrate: 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, rms_ref, qt_ref, scale_t_ref = torch_rms_and_block_quant_forward(x, - weight, - round_scale=True) + q_ref, scale_ref, rms_ref, qt_ref, scale_t_ref = torch_rms_and_block_quant_forward( + x, weight, round_scale=True + ) - q, scale, rms, _, _ = triton_rms_norm_and_block_quant_forward(x, weight, - round_scale=True, - output_mode=0) + 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') + 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) + _, _, _, 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") - q, scale, rms, q_t, scale_t = triton_rms_norm_and_block_quant_forward(x, - weight, - round_scale=True, - output_mode=2) + 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_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(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): + 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' + device = "cuda:0" 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) + 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) + 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(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(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') + output_check(scale_t_ref.t(), scale_t, name="1.block.scale") 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__': + 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) @@ -321,5 +334,4 @@ def test_rms_norm_fp32_gemm_block_quant_forward(M=8192, N=256, K=2048, test_rmsnorm_and_smooth_quant(M=8192, N=4096, bench=False) test_rmsnorm_and_smooth_quant(M=4096, N=8192, bench=False) - test_rms_norm_fp32_gemm_block_quant_forward(M=8192 * 2, N=256, K=2048, - 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 2490daf..e633038 100644 --- a/tests/test_rearange.py +++ b/tests/test_rearange.py @@ -24,7 +24,7 @@ def torch_sort_chunks_by_index(x, scales, counts, indices): def test_sort_chunks_by_index(M=4096, N=4096, bench=False): dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" n_repeat = 100 x = torch.randn(M, N, dtype=dtype, device=device) @@ -41,31 +41,38 @@ def test_sort_chunks_by_index(M=4096, N=4096, bench=False): scale_chunks = torch.split(x_scales, split_size_list) data_ref, scale_ref = torch_sort_chunks_by_index( - x_q.view(torch.float8_e4m3fn), - x_scales, split_size_list, - sorted_indices_list) + 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) + 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') + output_check(data_ref.view(torch.float8_e4m3fn), data, name="data") + output_check(scale_ref, scale, name="scale") if bench: - benchmark_func(torch.split, x_q.view(torch.uint8), split_size_list, - n_repeat=n_repeat) + 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.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_sort_chunks_by_index, - x_q.view(torch.float8_e4m3fn), - x_scales, split_size_list, sorted_indices_list, - n_repeat=n_repeat) - benchmark_func(triton_sort_chunks_by_index, x_q, counts, indices, - scales=x_scales, n_repeat=n_repeat) - - -if __name__ == '__main__': + benchmark_func( + torch_sort_chunks_by_index, + x_q.view(torch.float8_e4m3fn), + x_scales, + split_size_list, + sorted_indices_list, + n_repeat=n_repeat, + ) + benchmark_func( + triton_sort_chunks_by_index, + x_q, + counts, + indices, + scales=x_scales, + n_repeat=n_repeat, + ) + + +if __name__ == "__main__": test_sort_chunks_by_index(M=4096, N=4096) diff --git a/tests/test_reduce.py b/tests/test_reduce.py index 82506f1..ae9e1e6 100644 --- a/tests/test_reduce.py +++ b/tests/test_reduce.py @@ -9,10 +9,12 @@ from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check -from linghe.utils.reduce import (triton_abs_max, - triton_batch_count_zero, - triton_norm, - triton_batch_norm) +from linghe.utils.reduce import ( + triton_abs_max, + triton_batch_count_zero, + triton_norm, + triton_batch_norm, +) def torch_sum(xs, ord=2, norm=True): @@ -28,28 +30,32 @@ def torch_sum(xs, ord=2, norm=True): 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') + 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') + 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)] + 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 @@ -60,60 +66,80 @@ def test_count_zero(M=4096, N=8192, k=32, bench=False): 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) + 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_norm(M=4096, N=8192, coef=1.0, bench=False): - x = torch.randn(M, N, dtype=torch.float32, device='cuda:0') * 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') + 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') + output_check(sum_ref, sums, "l1_norm") if bench: ref_bytes = M * N * 4 n_repeat = 100 - ref_time = benchmark_func(lambda x: x.norm(p=2), x, n_repeat=n_repeat, - ref_bytes=ref_bytes) - benchmark_func(triton_norm, x, ord=2, norm=True, scalar=True, - n_repeat=n_repeat, - ref_bytes=ref_bytes, ref_time=ref_time) + ref_time = benchmark_func( + lambda x: x.norm(p=2), x, n_repeat=n_repeat, ref_bytes=ref_bytes + ) + benchmark_func( + triton_norm, + x, + ord=2, + norm=True, + scalar=True, + n_repeat=n_repeat, + ref_bytes=ref_bytes, + ref_time=ref_time, + ) def test_batch_norm(M=4096, N=8192, k=32, coef=1.0, bench=False): - 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)] + 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') + 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') + 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') + output_check(sum_ref, sums, "inf_norm") if bench: ref_bytes = sum([x.numel() for x in xs]) * 4 n_repeat = 100 - ref_time = benchmark_func(torch_sum, xs, 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) + ref_time = benchmark_func(torch_sum, xs, 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__': +if __name__ == "__main__": test_triton_abs_max(M=4096, N=4096, bench=False) test_count_zero(M=4096, N=8192, k=32, bench=False) test_norm(M=100000, N=8192, bench=False) diff --git a/tests/test_rope.py b/tests/test_rope.py index 4ed8994..6f9f127 100644 --- a/tests/test_rope.py +++ b/tests/test_rope.py @@ -8,20 +8,22 @@ from linghe.facade.rope import qk_norm_half_rope from linghe.tools.benchmark import benchmark_func 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) +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) @@ -33,16 +35,17 @@ def apply_rotary_pos_emb(q, k, cos, sin, position_ids): cos = cos[:, 0, 0][position_ids][:, :, None] sin = sin[:, 0, 0][position_ids][:, :, None] else: - raise ValueError('unsupported ndim=3') + 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 @@ -57,11 +60,12 @@ def torch_half_rope(q, k, freqs, transposed=True): 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.to(dtype), ko.to(dtype) @@ -88,25 +92,20 @@ def torch_mla_rope(q, kv, k_pos_emb, freqs, mscale=1.0, transpose=False): 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 - ) + q_no_pe, q_pos_emb = torch.split(q, [128, 64], dim=-1) - k_no_pe, value = torch.split( - kv, [128, 128], 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) + 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 = 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) + 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) @@ -124,8 +123,9 @@ def torch_mla_rope(q, kv, k_pos_emb, freqs, mscale=1.0, transpose=False): 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): +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) @@ -146,13 +146,9 @@ def torch_varlen_mla_rope(qs, kvs, k_pos_embs, freqs, lengths, mscale=1.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 - ) + 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 - ) + k_no_pe, value = torch.split(kv01[j], [128, 128], dim=-1) cos = freqs.cos().to(dtype) * mscale sin = freqs.sin().to(dtype) * mscale @@ -160,17 +156,20 @@ def torch_varlen_mla_rope(qs, kvs, k_pos_embs, freqs, lengths, mscale=1.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] + 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) + [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) + [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) + 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) @@ -189,9 +188,18 @@ def torch_varlen_mla_rope(qs, kvs, k_pos_embs, freqs, lengths, mscale=1.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): +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: @@ -230,9 +238,21 @@ def torch_qk_norm_and_half_rope(qkv, qw, kw, freqs, H=32, return q.to(dtype), k.to(dtype), v.to(dtype) -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): +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) @@ -254,15 +274,20 @@ def torch_varlen_qk_norm_and_half_rope(qkvs, qw, kw, freqs, lengths, H=32, h=4, 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') + 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) + 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) @@ -274,11 +299,11 @@ def torch_varlen_qk_norm_and_half_rope(qkvs, qw, kw, freqs, lengths, H=32, h=4, return qoss, koss, voss -def test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, - transposed=True, - bench=False): +def test_half_rope( + B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, transposed=True, bench=False +): dtype = torch.float32 - device = 'cuda:0' + 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) @@ -286,8 +311,8 @@ def test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, 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, name='q') - output_check(k_ref, ko, name='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) @@ -299,47 +324,67 @@ def test_half_rope(B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, 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, name='dq') - output_check(dk_ref, dk, name='dk', rtol=0.05, atol=0.1) + 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) - - -def test_qk_norm_and_half_rope(B=2, L=4096, H=32, h=8, D=128, - rope_theta=10000.0, - eps=1e-6, - interleaved=True, - transposed=True, - silu=False, - bench=False): + benchmark_func( + 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, + eps=1e-6, + interleaved=True, + transposed=True, + silu=False, + bench=False, +): 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) 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) + 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_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, 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) @@ -347,24 +392,38 @@ def test_qk_norm_and_half_rope(B=2, L=4096, H=32, h=8, D=128, 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) + 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 @@ -379,36 +438,63 @@ def test_qk_norm_and_half_rope(B=2, L=4096, H=32, h=8, D=128, 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') + output_check(dqkv_ref, dqkv, name="dqkv") + output_check(dqw_ref, dqw, name="dqw") + output_check(dkw_ref, dkw, name="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, - silu=silu, - 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, - transposed=transposed, interleaved=interleaved, - silu=silu, - ref_bytes=L * B * (H + 2 * h) * D * 6, - n_profile=0) - - -def test_varlen_qk_norm_and_half_rope(lengths=[2048, 2048], H=32, h=4, dim=128, - rope_theta=10000.0, silu=False, - interleaved=True, - bench=False, cp_size=1, cp_rank=0): + benchmark_func( + 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_func( + 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, + ) + + +def test_varlen_qk_norm_and_half_rope( + lengths=[2048, 2048], + H=32, + h=4, + dim=128, + rope_theta=10000.0, + silu=False, + interleaved=True, + bench=False, + cp_size=1, + cp_rank=0, +): dtype = torch.bfloat16 - device = 'cuda:0' + 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) + 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) @@ -418,21 +504,24 @@ def test_varlen_qk_norm_and_half_rope(lengths=[2048, 2048], H=32, h=4, dim=128, 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) + 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, 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) @@ -442,66 +531,96 @@ def test_varlen_qk_norm_and_half_rope(lengths=[2048, 2048], H=32, h=4, dim=128, 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=5.0, rtol=0.02) - output_check(dkw_ref, dkw.to(dtype), name='dkw', atol=5.0, rtol=0.02) + 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=5.0, rtol=0.02) + output_check(dkw_ref, dkw.to(dtype), name="dkw", atol=5.0, rtol=0.02) if bench: lbh = sum(lengths) // cp_size * H - benchmark_func(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_func(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) - - -def test_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=False, - bench=False): + benchmark_func( + 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_func( + 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, + ) + + +def test_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=False, bench=False): 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_() + 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] @@ -517,8 +636,9 @@ def test_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=False, mscale = 1.0 - q_ref, k_ref, v_ref = torch_mla_rope(q, kv, k_pos_emb, freqs, - mscale=mscale, transpose=transpose) + 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) @@ -526,189 +646,375 @@ def test_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=False, 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') + 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") if bench: lbh = L * B * H - benchmark_func(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_func(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) - - -def test_varlen_mla_rope(lengths=[2048, 2048], H=32, rope_theta=10000.0, - bench=False, cp_size=1, cp_rank=0): + benchmark_func( + 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_func( + 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, + ) + + +def test_varlen_mla_rope( + lengths=[2048, 2048], H=32, rope_theta=10000.0, bench=False, cp_size=1, cp_rank=0 +): dtype = torch.bfloat16 - 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_() + 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) + 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_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, 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) + 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) if bench: lbh = sum(lengths) // cp_size * H - benchmark_func(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_func(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) - - -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=2, L=4096, H=16, h=16, D=128, - rope_theta=10000.0, interleaved=True, - transposed=True, silu=True, bench=False) - test_qk_norm_and_half_rope(B=2, L=4096, H=16, h=16, D=128, - rope_theta=10000.0, interleaved=True, - transposed=True, silu=False, 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=False, silu=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=False, silu=False, bench=False) - test_qk_norm_and_half_rope(B=4, L=4096, H=32, h=4, D=128, - rope_theta=10000.0, interleaved=False, - transposed=True, silu=True, bench=False) - test_qk_norm_and_half_rope(B=4, L=4096, H=24, h=6, D=128, - rope_theta=10000.0, interleaved=True, - transposed=True, silu=False, bench=False) - test_qk_norm_and_half_rope(B=4, L=4096, H=32, h=32, D=128, - rope_theta=10000.0, interleaved=False, - transposed=False, silu=True, bench=False) - test_qk_norm_and_half_rope(B=1, L=4096, H=32, h=32, D=128, - rope_theta=10000.0, interleaved=False, - transposed=False, silu=False, bench=False) - test_varlen_qk_norm_and_half_rope(lengths=[2048], H=24, h=6, dim=128, - rope_theta=10000.0, silu=False, - interleaved=True, cp_size=1, cp_rank=0, - bench=False) - test_varlen_qk_norm_and_half_rope(lengths=[1024, 4096, 4096, 568], H=32, - h=4, dim=128, rope_theta=10000.0, - silu=False, - interleaved=True, cp_size=1, cp_rank=0, - bench=False) - test_varlen_qk_norm_and_half_rope(lengths=[2048, 4096, 4096], H=32, h=4, - dim=128, rope_theta=10000.0, silu=False, - interleaved=True, cp_size=1, cp_rank=0, - bench=False) - test_varlen_qk_norm_and_half_rope(lengths=[2048, 3072, 4096], H=32, h=4, - dim=128, rope_theta=10000.0, silu=False, - interleaved=True, cp_size=4, cp_rank=0, - bench=False) - test_varlen_qk_norm_and_half_rope(lengths=[2048, 4096, 4096], H=32, h=4, - dim=128, rope_theta=10000.0, silu=True, - interleaved=False, cp_size=4, cp_rank=0, - bench=False) - test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=False, - bench=False) - test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=True, - bench=False) - test_varlen_mla_rope(lengths=[8192], H=64, rope_theta=10000.0, cp_size=1, - cp_rank=0, - bench=False) - test_varlen_mla_rope(lengths=[4096, 4096], H=16, rope_theta=10000.0, - cp_size=1, cp_rank=0, - bench=False) - test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, - rope_theta=10000.0, cp_size=4, cp_rank=0, - bench=False) - test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, - rope_theta=10000.0, cp_size=4, cp_rank=1, - bench=False) - test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, - rope_theta=10000.0, cp_size=4, cp_rank=2, - bench=False) - test_varlen_mla_rope(lengths=[4096 * 4, 2048 * 4, 2048 * 4], H=32, - rope_theta=10000.0, cp_size=4, cp_rank=3, - bench=False) + benchmark_func( + 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_func( + 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, + ) + + +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=2, + L=4096, + H=16, + h=16, + D=128, + rope_theta=10000.0, + interleaved=True, + transposed=True, + silu=True, + bench=False, + ) + test_qk_norm_and_half_rope( + B=2, + L=4096, + H=16, + h=16, + D=128, + rope_theta=10000.0, + interleaved=True, + transposed=True, + silu=False, + 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=False, + silu=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=False, + silu=False, + bench=False, + ) + test_qk_norm_and_half_rope( + B=4, + L=4096, + H=32, + h=4, + D=128, + rope_theta=10000.0, + interleaved=False, + transposed=True, + silu=True, + bench=False, + ) + test_qk_norm_and_half_rope( + B=4, + L=4096, + H=24, + h=6, + D=128, + rope_theta=10000.0, + interleaved=True, + transposed=True, + silu=False, + bench=False, + ) + test_qk_norm_and_half_rope( + B=4, + L=4096, + H=32, + h=32, + D=128, + rope_theta=10000.0, + interleaved=False, + transposed=False, + silu=True, + bench=False, + ) + test_qk_norm_and_half_rope( + B=1, + L=4096, + H=32, + h=32, + D=128, + rope_theta=10000.0, + interleaved=False, + transposed=False, + silu=False, + bench=False, + ) + test_varlen_qk_norm_and_half_rope( + lengths=[2048], + H=24, + h=6, + dim=128, + rope_theta=10000.0, + silu=False, + interleaved=True, + cp_size=1, + cp_rank=0, + bench=False, + ) + test_varlen_qk_norm_and_half_rope( + lengths=[1024, 4096, 4096, 568], + H=32, + h=4, + dim=128, + rope_theta=10000.0, + silu=False, + interleaved=True, + cp_size=1, + cp_rank=0, + bench=False, + ) + test_varlen_qk_norm_and_half_rope( + lengths=[2048, 4096, 4096], + H=32, + h=4, + dim=128, + rope_theta=10000.0, + silu=False, + interleaved=True, + cp_size=1, + cp_rank=0, + bench=False, + ) + test_varlen_qk_norm_and_half_rope( + lengths=[2048, 3072, 4096], + H=32, + h=4, + dim=128, + rope_theta=10000.0, + silu=False, + interleaved=True, + cp_size=4, + cp_rank=0, + bench=False, + ) + test_varlen_qk_norm_and_half_rope( + lengths=[2048, 4096, 4096], + H=32, + h=4, + dim=128, + rope_theta=10000.0, + silu=True, + interleaved=False, + cp_size=4, + cp_rank=0, + bench=False, + ) + test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=False, bench=False) + test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=True, bench=False) + test_varlen_mla_rope( + lengths=[8192], H=64, rope_theta=10000.0, cp_size=1, cp_rank=0, bench=False + ) + test_varlen_mla_rope( + lengths=[4096, 4096], + H=16, + rope_theta=10000.0, + cp_size=1, + cp_rank=0, + bench=False, + ) + test_varlen_mla_rope( + lengths=[4096 * 4, 2048 * 4, 2048 * 4], + H=32, + rope_theta=10000.0, + cp_size=4, + cp_rank=0, + bench=False, + ) + test_varlen_mla_rope( + lengths=[4096 * 4, 2048 * 4, 2048 * 4], + H=32, + rope_theta=10000.0, + cp_size=4, + cp_rank=1, + bench=False, + ) + test_varlen_mla_rope( + lengths=[4096 * 4, 2048 * 4, 2048 * 4], + H=32, + rope_theta=10000.0, + cp_size=4, + cp_rank=2, + bench=False, + ) + test_varlen_mla_rope( + lengths=[4096 * 4, 2048 * 4, 2048 * 4], + H=32, + rope_theta=10000.0, + cp_size=4, + cp_rank=3, + bench=False, + ) diff --git a/tests/test_scatter.py b/tests/test_scatter.py index 89ff9b7..9c71505 100644 --- a/tests/test_scatter.py +++ b/tests/test_scatter.py @@ -8,10 +8,7 @@ from linghe.tools.benchmark import benchmark_func 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 - ) - +from linghe.utils.scatter import triton_scatter_add, triton_unpermute_with_mask_map # os.environ["CUDA_LAUNCH_BLOCKING"] = "1" @@ -29,11 +26,12 @@ def torch_scatter_add(x, outputs, indices, weights): def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): 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) @@ -43,23 +41,30 @@ 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_scatter_add(x, outputs.clone(), indices, None) - unpermuted_prob = probs.T.contiguous().masked_select( - mask_map.T.contiguous()) + 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') + 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(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__': + 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, bench=False) test_scatter(M=2467, N=4096, n_experts=32, topk=2, bias=-0.1, bench=False) test_scatter(M=2467, N=1536, n_experts=32, topk=2, bias=-0.1, bench=False) diff --git a/tests/test_silu.py b/tests/test_silu.py index 3713f45..00ac2af 100644 --- a/tests/test_silu.py +++ b/tests/test_silu.py @@ -10,19 +10,19 @@ 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 (torch_smooth_quant, - torch_group_quant) +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 torch_smooth_quant, torch_group_quant from linghe.tools.check import output_check @@ -60,8 +60,9 @@ def torch_silu_and_smooth_quant_forward(x, smooth_scale=None, round_scale=True): y = torch.sigmoid(x1) * x1 * x2 # smooth - y_q, y_scale, x_maxs = torch_smooth_quant(y, smooth_scale, reverse=False, - round_scale=round_scale) + 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 @@ -83,21 +84,29 @@ 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_smooth_quant_backward( + grad, + x, + smooth_scale=None, + transpose_smooth_scale=None, + round_scale=True, + reverse=True, +): grad = grad.float() x = x.float().detach().clone().requires_grad_() y = torch_silu(x) 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, 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, + ) return q, dx_scale, yt_q, yt_scale @@ -114,12 +123,9 @@ 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_batch_weighted_silu_and_smooth_quant_forward( + xs, weight, counts, smooth_scales=None, round_scale=True, reverse=False +): counts = counts.tolist() N = xs.shape[1] if sum(counts) == 0: @@ -138,10 +144,11 @@ def torch_batch_weighted_silu_and_smooth_quant_forward(xs, weight, 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]) + q, scale, ms = torch_smooth_quant( + y, smooth_scales[i], reverse=reverse, round_scale=round_scale + ) qs.append(q) scales.append(scale) maxs.append(ms) @@ -153,9 +160,9 @@ def torch_batch_weighted_silu_and_smooth_quant_forward(xs, weight, return qs, scales, maxs -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, round_scale=True +): counts = counts.tolist() N = xs.shape[1] if sum(counts) == 0: @@ -175,8 +182,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]) q, scale = torch_group_quant(y, round_scale=round_scale) qt, qtscale = torch_group_quant(y.t(), round_scale=round_scale) qs.append(q) @@ -192,12 +199,16 @@ 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_smooth_quant_backward( + grad_output, + x, + weight, + counts, + smooth_scales=None, + transpose_smooth_scale=None, + round_scale=True, + reverse=False, +): if sum(counts) == 0: device = x.device N = x.shape[1] @@ -205,8 +216,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() @@ -222,18 +232,18 @@ def torch_batch_weighted_silu_and_smooth_quant_backward(grad_output, x, weight, 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, 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] 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, dx_max = torch_smooth_quant( + dxt, dxt_s, reverse=reverse, round_scale=round_scale + ) qs.append(q) scales.append(scale) @@ -247,9 +257,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, round_scale=True +): if sum(counts) == 0: device = x.device N = x.shape[1] @@ -257,8 +267,7 @@ 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((0,), 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() @@ -272,9 +281,8 @@ def torch_batch_weighted_silu_and_block_quant_backward(grad_output, x, weight, 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)) @@ -288,312 +296,411 @@ 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, asm=False, coef=1.0, bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') + x = torch.randn((M, N), 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 // 2), dtype=torch.bfloat16, - device='cuda:0') + weight = torch.randn((M, 1), dtype=torch.float32, device="cuda:0") + grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, device="cuda:0") ref_y = torch_weighted_silu(x, weight) y = triton_weighted_silu_forward(x, weight, asm=asm) - output_check(ref_y, y, 'y') + 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', rtol=3e-3, atol=3e-3) + output_check(dx_ref, dx, "dx") + output_check(dw_ref, dw, "dw", rtol=3e-3, atol=3e-3) if bench: - benchmark_func(triton_weighted_silu_forward, x, weight, asm=asm, - 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, coef=1.0, grad_coef=1.0, - bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') + benchmark_func( + triton_weighted_silu_forward, + x, + weight, + asm=asm, + 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, coef=1.0, grad_coef=1.0, bench=False): + x = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") x = (x * coef).clone().detach().requires_grad_() - grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, - device='cuda:0') * grad_coef - 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') + grad_output = ( + torch.randn((M, N // 2), dtype=torch.bfloat16, device="cuda:0") * grad_coef + ) + 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" + ) round_scale = False - y_q_ref, y_scale_ref, y_maxs_ref = torch_silu_and_smooth_quant_forward(x, - smooth_scale=smooth_scale, - round_scale=round_scale) - y_q, y_scale, y_maxs = triton_silu_and_smooth_quant_forward(x, - smooth_scale=smooth_scale, - round_scale=round_scale, - calibrate=True) - output_check(y_q_ref, y_q, 'smooth.y_q', rtol=0.125) - 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) + y_q_ref, y_scale_ref, y_maxs_ref = torch_silu_and_smooth_quant_forward( + x, smooth_scale=smooth_scale, round_scale=round_scale + ) + y_q, y_scale, y_maxs = triton_silu_and_smooth_quant_forward( + x, smooth_scale=smooth_scale, round_scale=round_scale, calibrate=True + ) + output_check(y_q_ref, y_q, "smooth.y_q", rtol=0.125) + 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, + ) + ) 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) + 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') + 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") 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, coef=1.0, grad_coef=1.0, - bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device='cuda:0') + 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, coef=1.0, grad_coef=1.0, bench=False): + x = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") x = (x * coef).clone().detach().requires_grad_() - grad_output = torch.randn((M, N // 2), dtype=torch.bfloat16, - device='cuda:0') * grad_coef + grad_output = ( + torch.randn((M, N // 2), 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=round_scale) - - y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward(x, - round_scale=round_scale, - 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, - 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, - 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) + x, round_scale=round_scale + ) + + y_q, y_scale, yt_q, yt_scale = triton_silu_and_block_quant_forward( + x, round_scale=round_scale, 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, 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, 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) + ) dx_q, dx_scale, dxt_q, dxt_scale = triton_silu_and_block_quant_backward( - grad_output, x, - round_scale=round_scale) - 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') + grad_output, x, round_scale=round_scale + ) + 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") if bench: - benchmark_func(triton_silu_and_block_quant_forward, x, - round_scale=round_scale, output_mode=0, - n_repeat=100, ref_bytes=M * N * 3) - benchmark_func(triton_silu_and_block_quant_forward, x, - round_scale=round_scale, output_mode=1, - n_repeat=100, ref_bytes=M * N * 3) - benchmark_func(triton_silu_and_block_quant_forward, x, - round_scale=round_scale, output_mode=2, - 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, - coef=1.0, - grad_coef=1.0, - bench=False): - 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) + benchmark_func( + triton_silu_and_block_quant_forward, + x, + round_scale=round_scale, + output_mode=0, + n_repeat=100, + ref_bytes=M * N * 3, + ) + benchmark_func( + triton_silu_and_block_quant_forward, + x, + round_scale=round_scale, + output_mode=1, + n_repeat=100, + ref_bytes=M * N * 3, + ) + benchmark_func( + triton_silu_and_block_quant_forward, + x, + round_scale=round_scale, + output_mode=2, + 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, coef=1.0, grad_coef=1.0, bench=False +): + 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') * coef - 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 - - grad_output = torch.randn((bs, N // 2), dtype=torch.bfloat16, - device='cuda:0') * grad_coef - 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 + x = torch.randn((bs, 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 // 2), dtype=torch.float32, device="cuda:0") * 10 + ) + + grad_output = ( + torch.randn((bs, N // 2), dtype=torch.bfloat16, device="cuda:0") * grad_coef + ) + 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 + ) round_scale = True rtol = 2 if round_scale else 0.125 - x_q_ref, x_scale_ref, x_max_ref = torch_batch_weighted_silu_and_smooth_quant_forward( + x_q_ref, x_scale_ref, x_max_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_scales=smooth_scales, + smooth_scale=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, 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') + 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") 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, - 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) + 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, 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, N), dtype=torch.bfloat16, - device='cuda:0') * coef + x = torch.randn((bs, 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') + weight = torch.randn((bs, 1), dtype=torch.float32, device="cuda:0") - grad_output = torch.randn((bs, N // 2), dtype=torch.bfloat16, - device='cuda:0') * grad_coef + grad_output = ( + torch.randn((bs, N // 2), 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, - round_scale=round_scale) + x_q_ref, x_scale_ref, xt_q_ref, xt_scale_ref = ( + torch_batch_weighted_silu_and_block_quant_forward( + x, weight, counts, 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, - round_scale=round_scale, - output_mode=2) + x, weight, counts, count_list, round_scale=round_scale, 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') + 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") x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_block_quant_forward( - x, - weight, - counts, - count_list, - 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, weight, counts, count_list, 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, - round_scale=round_scale, - 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) - 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, dx, 'block.dx', rtol=rtol) - output_check(dx_scale_ref, dx_scale.view(-1), 'block.dx_scale') + x, weight, counts, count_list, round_scale=round_scale, 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 + ) + ) + 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, 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') + 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") 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=1, n_repeat=100, - ref_bytes=n_experts * M * N * 3, 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__': + 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=1, + n_repeat=100, + ref_bytes=n_experts * M * N * 3, + 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=4096, coef=1.0, asm=False, bench=False) test_weighted_silu(M=16384, N=4096, coef=1.0, asm=True, bench=False) test_weighted_silu(M=8192, N=1536, bench=False) @@ -608,24 +715,27 @@ def test_triton_batch_weighted_silu_and_block_quant(M=1024, N=4096, test_silu_and_block_quant(M=8192, N=4096, bench=False) test_silu_and_block_quant(M=16384, N=1536, bench=False) test_silu_and_block_quant(M=4096, N=1536 * 8, bench=False) - test_silu_and_block_quant(M=4096, N=1536 * 8, coef=100.0, grad_coef=100.0, - bench=False) - test_silu_and_block_quant(M=4096, N=1536 * 8, coef=0.0, grad_coef=0.0, - 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_smooth_quant(M=2048, N=2048, - n_experts=32, bench=False) - - test_triton_batch_weighted_silu_and_block_quant(M=0, N=1536, n_experts=32, - bench=False) - test_triton_batch_weighted_silu_and_block_quant(M=2048, N=8192, - n_experts=32, bench=False) - test_triton_batch_weighted_silu_and_block_quant(M=12080, N=1536, - n_experts=32, coef=100.0, - grad_coef=100.0, - bench=False) - test_triton_batch_weighted_silu_and_block_quant(M=12080, N=1536, - n_experts=32, coef=0.0, - grad_coef=0.0, bench=False) + test_silu_and_block_quant( + M=4096, N=1536 * 8, coef=100.0, grad_coef=100.0, bench=False + ) + test_silu_and_block_quant(M=4096, N=1536 * 8, coef=0.0, grad_coef=0.0, 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_smooth_quant( + M=2048, N=2048, n_experts=32, bench=False + ) + + test_triton_batch_weighted_silu_and_block_quant( + M=0, N=1536, n_experts=32, bench=False + ) + test_triton_batch_weighted_silu_and_block_quant( + M=2048, N=8192, n_experts=32, bench=False + ) + test_triton_batch_weighted_silu_and_block_quant( + M=12080, N=1536, n_experts=32, coef=100.0, grad_coef=100.0, bench=False + ) + test_triton_batch_weighted_silu_and_block_quant( + M=12080, N=1536, n_experts=32, coef=0.0, grad_coef=0.0, bench=False + ) diff --git a/tests/test_smooth_quant.py b/tests/test_smooth_quant.py index 1e0d45d..a16c826 100644 --- a/tests/test_smooth_quant.py +++ b/tests/test_smooth_quant.py @@ -6,16 +6,16 @@ import torch from linghe.facade.smooth_quant_linear import SmoothQuantLinear -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.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.check import output_check -from linghe.tools.util import (torch_make_indices, - torch_smooth_quant, - round_up) +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): @@ -35,11 +35,18 @@ def torch_split_smooth_quant(x_split, smooth_scales, round_scale=False): return x_qs, x_scales, x_maxs -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 @@ -47,7 +54,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 @@ -56,32 +63,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 @@ -96,154 +112,176 @@ def triton_split_smooth_quant(x_split, smooth_scales): def test_triton_smooth_quant(M=4096, N=4096, bench=False): - device = 'cuda:0' + device = "cuda:0" x = torch.randn((M, N), dtype=torch.bfloat16, device=device) - smooth_scale = torch.randn((N,), device=device, - dtype=torch.float32).abs() + 1.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, x_maxs_ref = torch_smooth_quant(x, smooth_scale, - reverse=False, - round_scale=round_scale) - - x_q, x_scale, x_maxs = triton_smooth_quant(x, smooth_scale, - reverse=False, - round_scale=round_scale, - calibrate=True) - output_check(x_q_ref, x_q, - 'triton_smooth_quant.data', rtol=rtol) - 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) + x_q_ref, scales_ref, x_maxs_ref = torch_smooth_quant( + x, smooth_scale, reverse=False, round_scale=round_scale + ) + x_q, x_scale, x_maxs = triton_smooth_quant( + x, smooth_scale, reverse=False, round_scale=round_scale, calibrate=True + ) + output_check(x_q_ref, x_q, "triton_smooth_quant.data", rtol=rtol) + output_check(scales_ref, x_scale, "triton_smooth_quant.scale") + output_check(x_maxs_ref, x_maxs, "triton_smooth_quant.x_maxs") -def test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, - size=16384): - device = 'cuda:0' + 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" 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') + 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' + 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, maxs_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') + 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' + 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" 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, 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, + ) 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], @@ -252,56 +290,77 @@ def test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, # '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' +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) + 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_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 = 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_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, 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) + torch.float8_e4m3fn + ) x_scale_ref = torch.cat(x_scale_ref, 0) 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') - output_check(x_maxs_ref.float(), x_maxs.float(), - 'triton_batch_smooth_quant.maxs') + 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" + ) + 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) + 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, + ) def test_smooth_quant_linear(M=8192, N=1024, K=2048): dtype = torch.bfloat16 - device = 'cuda:0' + 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) @@ -310,18 +369,18 @@ def test_smooth_quant_linear(M=8192, N=1024, K=2048): y_ref = x @ w.t() y = linear(x) - output_check(y_ref, y, name='y') + output_check(y_ref, y, name="y") dx_ref = dy @ w dw_ref = dy.t() @ x y.backward(dy) dw = linear.weight.grad dx = x.grad - output_check(dx_ref, dx, name='dx') - output_check(dw_ref, dw, name='dw') + output_check(dx_ref, dx, name="dx") + output_check(dw_ref, dw, name="dw") -if __name__ == '__main__': +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) @@ -330,27 +389,21 @@ def test_smooth_quant_linear(M=8192, N=1024, K=2048): 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_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_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) diff --git a/tests/test_topk.py b/tests/test_topk.py index b5d5c7d..6eeaddd 100644 --- a/tests/test_topk.py +++ b/tests/test_topk.py @@ -8,25 +8,28 @@ 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) +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, + 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) + 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) @@ -39,45 +42,60 @@ def group_limited_topk( .reshape(num_tokens, -1) ) - masked_scores = scores.masked_fill(~score_mask.bool(), float('-inf')) + 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): +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_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 + 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() + 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 @@ -85,7 +103,7 @@ def torch_group_topk_score(logits, expert_bias=None, num_experts=256, topk=8, def test_topk(M=4096, B=1, N=256, k=8, equal=False, bench=False): dtype = torch.float32 - device = 'cuda:0' + device = "cuda:0" if B == 0: x = torch.randn(M, N, dtype=dtype, device=device) @@ -97,8 +115,9 @@ def test_topk(M=4096, B=1, N=256, k=8, equal=False, bench=False): x = x.requires_grad_() if equal: - xd = x.to(torch.float64) * (1 - torch.arange(0, N, device=device).to( - torch.float64) * 1e-12) + 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: @@ -111,32 +130,40 @@ def test_topk(M=4096, B=1, N=256, k=8, equal=False, bench=False): 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') + 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') + 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): + 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' + device = "cuda:0" x = torch.randn(M, N, dtype=dtype, device=device) @@ -149,63 +176,109 @@ def test_group_topk_score(M=4096, N=256, k=8, num_groups=32, group_topk=4, 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) + 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, 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') + 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__': + 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) + 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 6cbaa96..317f617 100644 --- a/tests/test_transpose.py +++ b/tests/test_transpose.py @@ -9,11 +9,13 @@ from linghe.tools.benchmark import benchmark_func 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 linghe.utils.transpose import ( + round_up, + triton_batch_transpose, + triton_batch_transpose_and_pad, + triton_transpose, + triton_transpose_and_pad, +) def torch_nd_transpose(x, dim0, dim1): @@ -32,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 @@ -45,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 @@ -55,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): @@ -67,43 +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, inner=True) - output_check(t_ref, t, '3d_transpose') + 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, inner=True) - output_check(t_ref, t, '3d_transpose_stride') + 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, inner=True) - output_check(t_ref, t, '4d_transpose') + 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') + 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, inner=True, 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) + 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): @@ -112,7 +125,7 @@ def test_transpose_and_pad(M=4095, N=4096, bench=False): # M, N, K = 4096, 4096, 4096 dtype = torch.bfloat16 - device = 'cuda:0' + device = "cuda:0" x = torch.randn(M, N, dtype=dtype, device=device) P = round_up(M, b=32) @@ -121,66 +134,78 @@ def test_transpose_and_pad(M=4095, N=4096, bench=False): 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, - 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) x_t_ref = torch.cat([x.view(-1) for x in x_t_ref]) - output_check(x_t_ref, xts, f'batch_transpose') + 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]) - output_check(x_t_ref, x_t, - f'batch_transpose_and_pad') + 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=2048, bench=False) diff --git a/tests/test_unary.py b/tests/test_unary.py index b0938bb..7065220 100644 --- a/tests/test_unary.py +++ b/tests/test_unary.py @@ -12,15 +12,12 @@ from linghe.utils.unary import triton_calculate_smooth_scale, triton_batch_clip -def torch_calculate_smooth_scale(x, min_value=1.0, smooth_coef=0.5, - round_scale=False): +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) + 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))) + weight_smooth_scales = torch.exp2(torch.ceil(torch.log2(weight_smooth_scales))) return weight_smooth_scales @@ -31,64 +28,80 @@ def torch_batch_clip(xs, clip_value): def test_calculate_smooth_scale(N=4096, bench=False): - x = torch.randn(N, dtype=torch.float32, device='cuda:0').abs() ** 3 + 0.1 + x = torch.randn(N, dtype=torch.float32, device="cuda:0").abs() ** 3 + 0.1 min_value = 0.0 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') + 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) - - -def test_batch_clip(M=2048, N=1024, k=1024, clip_value=1.0, inf=False, - bench=False): - 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)] + 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, + ) + + +def test_batch_clip(M=2048, N=1024, k=1024, clip_value=1.0, inf=False, bench=False): + 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') + 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') + 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 bench: 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_func(torch_batch_clip, xs3, clip_value, - ref_bytes=ref_bytes, - n_repeat=n_repeat, - n_warmup=0) + n_repeat = 1 # inplace update will speedup our triton op + ref_time = benchmark_func( + 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_func(triton_batch_clip, xs4, clip_value, - ref_bytes=ref_bytes, ref_time=ref_time, - n_repeat=n_repeat, - n_warmup=0) - - -if __name__ == '__main__': + benchmark_func( + triton_batch_clip, + xs4, + clip_value, + ref_bytes=ref_bytes, + ref_time=ref_time, + n_repeat=n_repeat, + n_warmup=0, + ) + + +if __name__ == "__main__": # test_calculate_smooth_scale(N=4096*32) # test_calculate_smooth_scale(N=4096*32-1897) # test_batch_clip(M=2048, N=8192, k=128, clip_value=0.1, bench=False) # test_batch_clip(M=2048, N=1024, k=128, clip_value=1.0, bench=False) - test_batch_clip(M=2048, N=1024, k=128, clip_value=100.0, inf=True, - bench=False) + test_batch_clip(M=2048, N=1024, k=128, clip_value=100.0, inf=True, bench=False) From 0994eadc76855fe6169fe1048c57dc2cf841a497 Mon Sep 17 00:00:00 2001 From: "liangchen.liangche" Date: Mon, 27 Apr 2026 11:33:27 +0800 Subject: [PATCH 08/11] sync utils --- linghe/utils/add.py | 180 ++- linghe/utils/emb.py | 110 +- linghe/utils/gate.py | 511 +++++++-- linghe/utils/gather.py | 815 ++++++++----- linghe/utils/loss.py | 1 - linghe/utils/norm.py | 305 ++++- linghe/utils/reduce.py | 18 +- linghe/utils/rope.py | 2262 +++++++++---------------------------- linghe/utils/scatter.py | 88 ++ linghe/utils/silu.py | 1264 ++++++++++++++++----- linghe/utils/topk.py | 209 +++- linghe/utils/transpose.py | 53 +- linghe/utils/unary.py | 52 +- 13 files changed, 3339 insertions(+), 2529 deletions(-) diff --git a/linghe/utils/add.py b/linghe/utils/add.py index dbfe7c1..6d774c9 100644 --- a/linghe/utils/add.py +++ b/linghe/utils/add.py @@ -3,6 +3,8 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +from typing import List + import torch import triton import triton.language as tl @@ -10,58 +12,28 @@ @triton.jit def inplace_add_kernel( - x_ptr, - y_ptr, - M, - N, - H: tl.constexpr, - W: tl.constexpr, - EVEN: tl.constexpr, - ACCUM: tl.constexpr, + x_ptr, y_ptr, N, B: 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, :] - ) + 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): @@ -76,17 +48,123 @@ def triton_inplace_add(x: torch.Tensor, y: torch.Tensor, accum: bool = True): updated x """ assert x.is_contiguous() and y.is_contiguous() - N = x.shape[-1] - M = x.numel() // N - # M, N = x.shape - H = 128 - W = 128 - EVEN = M % H == 0 and N % W == 0 + 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, accum, num_stages=num_stages, num_warps=num_warps + 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, + ) + return xs diff --git a/linghe/utils/emb.py b/linghe/utils/emb.py index af4edfa..756ddcd 100644 --- a/linghe/utils/emb.py +++ b/linghe/utils/emb.py @@ -6,6 +6,7 @@ import torch import triton import triton.language as tl +import math @triton.jit @@ -56,8 +57,10 @@ def atomic_embedding_backward_kernel( if T == 0: grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) - else: + 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), @@ -78,10 +81,15 @@ def triton_atomic_embedding_backward(y, x, g_ptr, dtype=torch.bfloat16): Returns: None """ - assert dtype in (torch.bfloat16, torch.float32) + assert dtype in (torch.float32, torch.bfloat16, torch.float16) shape = x.shape assert len(shape) == 2 - T = 0 if dtype == torch.float32 else 1 + 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) @@ -134,8 +142,10 @@ def sync_embedding_backward_kernel( if T == 0: grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) - else: + 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) @@ -165,8 +175,13 @@ def triton_sync_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): Returns: None """ - assert dtype in (torch.bfloat16, torch.float32) - T = 0 if dtype == torch.float32 else 1 + 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 @@ -207,7 +222,9 @@ def scan_and_count_split_kernel( sid = tl.program_id(axis=1) ns = tl.num_programs(1) - ids = tl.load(id_ptr + bid * L + sid * B + tl.arange(0, B)) + 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 @@ -266,35 +283,21 @@ def triton_scan_and_count(ids): shape = ids.shape device = ids.device assert len(shape) in (1, 2) + + BLOCK = 128 if len(shape) == 2: B, L = ids.shape - BLOCK = 256 - assert L % BLOCK == 0 - T = 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) else: - L = shape[0] - B = 1 - BLOCK = 256 - assert L % BLOCK == 0 - T = L // BLOCK - counts = torch.empty((L,), dtype=torch.int32, device=device) - unique_ids = torch.empty((L,), dtype=torch.int32, device=device) - unique_counts = torch.empty((T,), dtype=torch.int32, device=device) - accum_counts = torch.zeros((L + 1,), dtype=torch.int32, device=device) + 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 = 1 + num_warps = 4 grid = (B, T) scan_and_count_split_kernel[grid]( ids, @@ -308,7 +311,7 @@ def triton_scan_and_count(ids): ) num_stages = 3 - num_warps = 1 + num_warps = 4 grid = (B,) scan_and_count_merge_kernel[grid]( counts, @@ -323,6 +326,9 @@ def triton_scan_and_count(ids): ) accum_counts = torch.cumsum(accum_counts, -1) + if len(shape) == 1: + accum_counts = accum_counts.squeeze(0) + return accum_counts @@ -379,10 +385,11 @@ def embedding_backward_kernel( dim, B, L, - DIM: tl.constexpr, + 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: @@ -393,24 +400,33 @@ def embedding_backward_kernel( if T == 0: grad_ptr = g_ptr.to(tl.pointer_type(tl.float32)) - else: + 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) + 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 + tl.arange(0, DIM), - mask=tl.arange(0, DIM) < dim, + 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 + tl.arange(0, DIM), + grad_ptr + input_id * dim + cid * BLOCK + tl.arange(0, BLOCK), outputs, - mask=tl.arange(0, DIM) < dim, + mask=mask, ) @@ -424,8 +440,13 @@ def triton_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): Returns: None """ - assert dtype in (torch.bfloat16, torch.float32) - T = 0 if dtype == torch.float32 else 1 + 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 @@ -434,11 +455,12 @@ def triton_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): 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) + BLOCK = 512 + assert dim % BLOCK == 0 num_stages = 3 num_warps = 2 - grid = (B * L,) + grid = (B * L, dim // BLOCK) embedding_backward_kernel[grid]( grad_output, sorted_ids, @@ -450,7 +472,7 @@ def triton_embedding_backward(grad_output, x, g_ptr, dtype=torch.bfloat16): dim, B, L, - DIM, + BLOCK, T, num_stages=num_stages, num_warps=num_warps, diff --git a/linghe/utils/gate.py b/linghe/utils/gate.py index 57977b3..aaada50 100644 --- a/linghe/utils/gate.py +++ b/linghe/utils/gate.py @@ -1,15 +1,16 @@ import torch import triton import triton.language as tl +import math -# TOOD(nanxiao): opt performance @triton.jit def group_rms_norm_gate_forward_kernel( x_ptr, gate_ptr, weight_ptr, out_ptr, + stride_g, eps, bs, length, @@ -18,7 +19,7 @@ def group_rms_norm_gate_forward_kernel( D: tl.constexpr, GROUP_SIZE: tl.constexpr, SHARE: tl.constexpr, - TRANSPOSE: tl.constexpr, + NATIVE: tl.constexpr, ): pid = tl.program_id(axis=0) bid = pid // length @@ -34,39 +35,45 @@ def group_rms_norm_gate_forward_kernel( mask=tl.arange(0, D)[None, :] < d, ) - x_offs = ( - pid * 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) - if TRANSPOSE: - g_offs = ( + 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, :] ) - g = tl.load(gate_ptr + g_offs, mask=tl.arange(0, D)[None, :] < d).to(tl.float32) - else: - g = tl.load(gate_ptr + x_offs, mask=x_offs_mask).to(tl.float32) + + 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) - if TRANSPOSE: - tl.store(out_ptr + g_offs, x, mask=tl.arange(0, D)[None, :] < d) - else: - tl.store(out_ptr + x_offs, x, mask=x_offs_mask) + 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, - transpose=True, + x: torch.Tensor, gate: torch.Tensor, weight: torch.Tensor, eps=1e-6, group_size=4 ): """ norm and gate in linear attention @@ -76,29 +83,26 @@ def triton_group_rms_norm_gate_forward( 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 + layout: layout of x, should in {'bsd', 'sbd} Returns: - output tensor, [length, bs, dim] if transpose=True else [bs, length, dim] + output tensor, [length, bs, dim] """ - # row-wise read, row-wise write - if transpose: - length, bs, dim = gate.shape - else: - bs, length, dim = gate.shape + length, bs, dim = gate.shape + assert dim <= 8192 and triton.next_power_of_2(group_size) == group_size - assert x.is_contiguous() and gate.is_contiguous() and weight.is_contiguous() + 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 + 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) - 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) + out = torch.empty((length, bs, dim), device=device, dtype=x.dtype) grid = (bs * length,) group_rms_norm_gate_forward_kernel[grid]( @@ -106,6 +110,7 @@ def triton_group_rms_norm_gate_forward( gate, weight, out, + gate.stride(1), eps, bs, length, @@ -113,8 +118,8 @@ def triton_group_rms_norm_gate_forward( d, D, group_size, - share, - transpose, + SHARE, + NATIVE, num_stages=3, num_warps=4, ) @@ -133,13 +138,14 @@ def group_rms_norm_gate_backward_kernel( eps, bs, length, + stride_g, DIM: tl.constexpr, d: tl.constexpr, D: tl.constexpr, GROUP_SIZE: tl.constexpr, T: tl.constexpr, SHARE: tl.constexpr, - TRANSPOSE: tl.constexpr, + NATIVE: tl.constexpr, ): pid = tl.program_id(0) bid = pid * T // length @@ -153,28 +159,40 @@ def group_rms_norm_gate_backward_kernel( mask=tl.arange(0, D)[None, :] < d, ) - x_offs = ( - pid * DIM * T + tl.arange(0, GROUP_SIZE)[:, None] * d + tl.arange(0, D)[None, :] - ) - x_offs_mask = tl.arange(0, D)[None, :] < d - if TRANSPOSE: - offs = ( + 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, :] ) - offs_mask = tl.arange(0, D)[None, :] < d + + 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) - if TRANSPOSE: - g = tl.load(grad_output_ptr + offs, offs_mask).to(tl.float32) - gate = tl.load(gate_ptr + offs, offs_mask).to(tl.float32) - else: - g = tl.load(grad_output_ptr + x_offs, mask=x_offs_mask).to(tl.float32) - gate = tl.load(gate_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 @@ -188,14 +206,14 @@ def group_rms_norm_gate_backward_kernel( tl.store(dx_ptr + x_offs, dx, mask=x_offs_mask) dg = x * r * w * g * gate * (1 - gate) - if TRANSPOSE: - tl.store(dg_ptr + offs, dg, mask=offs_mask) - else: - tl.store(dg_ptr + x_offs, dg, mask=x_offs_mask) + tl.store(dg_ptr + offs, dg, mask=offs_mask) - x_offs += DIM - if TRANSPOSE: - offs += DIM * bs + 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) @@ -212,17 +230,16 @@ def group_rms_norm_gate_backward_kernel( def triton_group_rms_norm_gate_backward( - grad_output, x, gate, weight, eps=1e-6, group_size=4, transpose=True + grad_output, x, gate, weight, eps=1e-6, group_size=4 ): - if transpose: - length, bs, dim = gate.shape - else: - bs, length, dim = gate.shape + 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 + SHARE = wd != dim # all groups share the same weight + NATIVE = x.size(0) == bs device = x.device dx = torch.empty_like(x) @@ -230,7 +247,7 @@ def triton_group_rms_norm_gate_backward( T = 8 g = (bs * length) // T - if share: + 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) @@ -248,15 +265,375 @@ def triton_group_rms_norm_gate_backward( eps, bs, length, + gate.stride(1), dim, d, D, group_size, T, - share, - transpose, + 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 b250205..87a01cb 100644 --- a/linghe/utils/gather.py +++ b/linghe/utils/gather.py @@ -10,6 +10,100 @@ import triton.language as tl +@triton.jit +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 @@ -219,7 +313,7 @@ def triton_make_row_id_map_and_index( @triton.jit -def index_select_kernel( +def permute_with_indices_kernel( x_ptr, out_ptr, scale_ptr, @@ -242,9 +336,9 @@ def index_select_kernel( 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] @@ -266,7 +360,7 @@ def triton_index_select(x, indices, scale=None, out=None, scale_out=None): T = triton.cdiv(E, sm) SCALE = scale is not None grid = (sm,) - index_select_kernel[grid]( + 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 @@ -462,10 +556,132 @@ def triton_permute_with_mask_map( @triton.jit -def batch_smooth_transpose_smooth_permute_kernel( +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, - scale_ptr, - oss_ptr, ss_ptr, index_ptr, count_ptr, @@ -476,7 +692,6 @@ def batch_smooth_transpose_smooth_permute_kernel( E: tl.constexpr, H: tl.constexpr, W: tl.constexpr, - SMOOTHED: tl.constexpr, ROUND: tl.constexpr, ): eid = tl.program_id(axis=0) @@ -490,9 +705,6 @@ def batch_smooth_transpose_smooth_permute_kernel( loop = tl.cdiv(pad, H) 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) @@ -504,11 +716,7 @@ def batch_smooth_transpose_smooth_permute_kernel( 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 = x * smooth_scale x_max = tl.maximum(tl.abs(x), x_max) scale = tl.maximum(tl.max(x_max, 0) / 448.0, 1e-30) @@ -529,11 +737,8 @@ def batch_smooth_transpose_smooth_permute_kernel( 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 + + 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 @@ -541,8 +746,6 @@ def batch_smooth_transpose_smooth_permute_kernel( def triton_batch_transpose_smooth_permute_with_indices( x, - scale, - org_smooth_scale, smooth_scales, indices, token_count_per_expert, @@ -552,13 +755,11 @@ def triton_batch_transpose_smooth_permute_with_indices( 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 @@ -580,20 +781,19 @@ def triton_batch_transpose_smooth_permute_with_indices( 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, @@ -604,7 +804,6 @@ def triton_batch_transpose_smooth_permute_with_indices( n_expert, H, W, - smoothed, round_scale, num_stages=3, num_warps=8, @@ -613,115 +812,7 @@ def triton_batch_transpose_smooth_permute_with_indices( @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).to(tl.int64) - 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] - """ - assert grads.is_contiguous() - M, N = grads.shape - n_expert, n = smooth_scales.shape - assert N == n, f"{N=} {n=}" - assert triton.next_power_of_2(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( +def batch_smooth_fused_permute_with_indices_kernel( grads_data_ptr, grads_scale_ptr, q_ptr, @@ -731,16 +822,13 @@ def smooth_permute_with_indices_kernel( accum_ptr, index_ptr, N: tl.constexpr, - hs: tl.constexpr, REVERSE: tl.constexpr, ROUND: tl.constexpr, - GROUP: 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 @@ -751,12 +839,9 @@ def smooth_permute_with_indices_kernel( 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)) @@ -773,9 +858,10 @@ def smooth_permute_with_indices_kernel( tl.store(q_ptr + i * N + tl.arange(0, N), xq) -def triton_smooth_permute_with_indices( +def triton_batch_smooth_fused_permute_with_indices( grad_data, grad_scale, + org_smooth_scales, smooth_scales, token_count_per_expert, indices, @@ -785,11 +871,12 @@ def triton_smooth_permute_with_indices( 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] @@ -807,8 +894,7 @@ def triton_smooth_permute_with_indices( 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 @@ -818,8 +904,9 @@ def triton_smooth_permute_with_indices( 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, @@ -829,10 +916,8 @@ def triton_smooth_permute_with_indices( accum_token_count, indices, N, - hs, reverse, round_scale, - group, num_stages=3, num_warps=16, ) @@ -840,121 +925,153 @@ def triton_smooth_permute_with_indices( @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, +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, - hs: tl.constexpr, - REVERSE: 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 - ) + 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) - 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 + 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 - x *= smooth_scale - x_max = tl.max(tl.abs(x)) + 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(x_max / 448.0, 1e-30) - if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + scale = tl.maximum(tl.max(x_max, 0) / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(quant_scale_ptr + index, scale, mask=mask) + tl.store(qs_ptr + eid * N + cid * W + tl.arange(0, W), scale) - 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) + 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_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, +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 inp.is_contiguous() - assert row_id_map.shape[1] == num_experts - assert triton.next_power_of_2(hidden_size) == hidden_size - 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 - ) - - sm = 128 - T = triton.cdiv(num_tokens, sm) - grid = (num_experts, sm) - smooth_permute_with_mask_map_kernel[grid]( - inp, - output, - row_id_map, + 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, - permuted_scale, - num_tokens, - T, - hidden_size, - hs, - reverse, + 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 output, permuted_scale + return x_q, x_scale @triton.jit @@ -1050,7 +1167,7 @@ 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 + select and quant, used in megatron 0.12 megatron flex backend Args: xs: [bs, dim] token_count_per_expert: [n_experts] @@ -1111,3 +1228,195 @@ def triton_batch_block_pad_permute_with_indices( ) 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 x_q, x_scale, xt_q, xt_scale, prob_output diff --git a/linghe/utils/loss.py b/linghe/utils/loss.py index cf63b15..0c3cb6c 100644 --- a/linghe/utils/loss.py +++ b/linghe/utils/loss.py @@ -343,7 +343,6 @@ def triton_parallel_softmax_cross_entropy_forward( num_stages=3, num_warps=2, ) - return loss, sum_exp, max_logit diff --git a/linghe/utils/norm.py b/linghe/utils/norm.py index 77133a8..e7b393e 100644 --- a/linghe/utils/norm.py +++ b/linghe/utils/norm.py @@ -509,15 +509,12 @@ def rms_norm_and_smooth_quant_forward_kernel( 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) @@ -525,21 +522,15 @@ def rms_norm_and_smooth_quant_forward_kernel( 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, :] 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) - if OUTPUT: - tl.store(rms_ptr + indices, rms, mask=indices < M) + 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: @@ -549,9 +540,6 @@ def rms_norm_and_smooth_quant_forward_kernel( 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( @@ -562,8 +550,6 @@ def triton_rms_norm_and_smooth_quant_forward( out=None, scale=None, rms=None, - calibrate=False, - output_rms=False, round_scale=False, ): """""" @@ -581,11 +567,7 @@ def triton_rms_norm_and_smooth_quant_forward( 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]( @@ -594,23 +576,292 @@ def triton_rms_norm_and_smooth_quant_forward( 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, ) - 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 diff --git a/linghe/utils/reduce.py b/linghe/utils/reduce.py index 40960c8..ccac83e 100644 --- a/linghe/utils/reduce.py +++ b/linghe/utils/reduce.py @@ -103,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 tl.range(t, flatten=True): + 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 @@ -155,7 +155,7 @@ def norm_kernel(input_ptr, tmp_ptr, m, B: tl.constexpr, ORD: tl.constexpr): tl.store(tmp_ptr + pid, sums) -def triton_norm(x, ord=2, norm=True, scalar=True): +def triton_norm(x, ord=2, norm=True, scalar=True, dtype=torch.float32): """ calculate norm. Args: @@ -176,7 +176,7 @@ def triton_norm(x, ord=2, norm=True, scalar=True): m = x.numel() B = 512 T = triton.cdiv(m, B) - tmp = torch.empty((T,), device=device, dtype=torch.float32) + 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: @@ -202,7 +202,7 @@ def batch_norm_kernel( ): tid = tl.program_id(axis=0) bid = tl.program_id(axis=1).to(tl.int64) - sm = tl.num_programs(axis=1) + T = tl.num_programs(axis=1) if HP: sums = tl.zeros((B,), dtype=tl.float64) else: @@ -213,7 +213,7 @@ def batch_norm_kernel( 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 * sm) + 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) @@ -233,7 +233,7 @@ def batch_norm_kernel( sums = tl.max(sums) else: sums = tl.sum(sums) - tl.store(tmp_ptr + tid * sm + bid, sums) + tl.store(tmp_ptr + tid * T + bid, sums) def triton_batch_norm(xs, ord=2, norm=True, scalar=True, high_precision=True): @@ -266,15 +266,15 @@ def triton_batch_norm(xs, ord=2, norm=True, scalar=True, high_precision=True): ) DT = 0 if dtype == torch.float32 else 1 - sm = 256 + T = 256 tensor_count = len(xs) tmp = torch.empty( - (tensor_count, sm), + (tensor_count, T), device=device, dtype=torch.float64 if high_precision else torch.float32, ) B = 128 - grid = (tensor_count, sm) + grid = (tensor_count, T) batch_norm_kernel[grid]( ptrs, sizes, tmp, DT, B, ord, high_precision, num_stages=2, num_warps=2 ) diff --git a/linghe/utils/rope.py b/linghe/utils/rope.py index 5f70afc..8931041 100644 --- a/linghe/utils/rope.py +++ b/linghe/utils/rope.py @@ -389,254 +389,6 @@ def triton_half_rope_backward(q_grad, k_grad, freqs, inplace=False, transposed=T @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, - 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)).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)).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, H) + tl.arange(0, H) // w * 2 - else: - row_offs = tl.arange(0, 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, :] - ).to(tl.float32) - q1 = tl.load( - q_ptr - + pid * B * stride - + i * stride - + D - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - else: - q0 = tl.load( - q_ptr - + i * L * stride - + pid * stride - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - q1 = tl.load( - q_ptr - + i * L * stride - + pid * stride - + D - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).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 - + i * L * H * DD - + D - + DD * tl.arange(0, H)[:, None] - + tl.arange(0, D)[None, :], - q1, - ) - - 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), - ) - 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)).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) - k_ptr = qkv_ptr + DD * w - else: - row_offs = tl.arange(0, 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, :] - ).to(tl.float32) - k1 = tl.load( - k_ptr - + pid * B * stride - + i * stride - + D - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - else: - k0 = tl.load( - k_ptr - + i * L * stride - + pid * stride - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - k1 = tl.load( - k_ptr - + i * L * stride - + pid * stride - + D - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).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 - tl.store( - ko_ptr - + pid * h * DD - + i * L * h * DD - + D - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :], - k1, - ) - - 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), - ) - 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, - ) - - if INTERLEAVED: - row_offs = tl.arange(0, h) * (w + 2) - v_ptr = qkv_ptr + DD * w + DD - else: - row_offs = tl.arange(0, 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, :] - ).to(tl.float32) - v1 = tl.load( - v_ptr - + pid * B * stride - + i * stride - + D - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - else: - v0 = tl.load( - v_ptr - + i * L * stride - + pid * stride - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - v1 = tl.load( - v_ptr - + i * L * stride - + pid * stride - + D - + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - if SILU: - v0 = v0 * tl.sigmoid(v0) - v1 = v1 * tl.sigmoid(v1) - - tl.store( - vo_ptr - + pid * h * DD - + i * L * h * DD - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :], - v0, - ) - tl.store( - vo_ptr - + pid * h * DD - + i * L * h * DD - + D - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :], - v1, - ) - - -@triton.jit -def compatible_qk_norm_and_half_rop_forward_kernel( qkv_ptr, q_norm_weight_ptr, k_norm_weight_ptr, @@ -983,52 +735,29 @@ def triton_qk_norm_and_half_rope_forward( H_p = triton.next_power_of_2(H) h_p = triton.next_power_of_2(h) - if H_p == H and h_p == h: - qk_norm_and_half_rope_forward_kernel[grid]( - qkv, - q_norm_weight, - k_norm_weight, - freqs, - qo, - ko, - vo, - B, - stride, - eps, - H, - h, - D // 2, - D // 4, - interleaved, - transposed, - silu, - num_stages=num_stages, - num_warps=num_warps, - ) - else: - compatible_qk_norm_and_half_rop_forward_kernel[grid]( - qkv, - q_norm_weight, - k_norm_weight, - freqs, - 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, - ) + qk_norm_and_half_rope_forward_kernel[grid]( + qkv, + q_norm_weight, + k_norm_weight, + freqs, + 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, + ) return qo, ko, vo @@ -1050,6 +779,8 @@ def qk_norm_and_half_rope_backward_kernel( eps, H: tl.constexpr, h: tl.constexpr, + H_p: tl.constexpr, + h_p: tl.constexpr, D: tl.constexpr, d: tl.constexpr, INTERLEAVED: tl.constexpr, @@ -1066,8 +797,8 @@ def qk_norm_and_half_rope_backward_kernel( sin = tl.sin(freqs) 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) + q_w0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)) + q_w1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)) dqw_0 = tl.zeros((D,), dtype=tl.float32) dqw_1 = tl.zeros((D,), dtype=tl.float32) @@ -1075,34 +806,40 @@ def qk_norm_and_half_rope_backward_kernel( 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 - else: - row_offs = tl.arange(0, H) + # 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_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, :] + + 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, :] + + 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, 2, d)), (0, 2, 1)), dim=2) + tl.flip(tl.permute(tl.reshape(gq_0, (H_p, 2, d)), (0, 2, 1)), dim=2) * signs, (0, 2, 1), ), - (H, D), + (H_p, D), ) gq_0 = gq_0 * cos + gq_r * sin @@ -1112,7 +849,8 @@ def qk_norm_and_half_rope_backward_kernel( + pid * B * stride + i * stride + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) q1 = tl.load( q_ptr @@ -1120,15 +858,18 @@ def qk_norm_and_half_rope_backward_kernel( + i * stride + D + DD * row_offs[:, None] - + tl.arange(0, D)[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, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) q1 = tl.load( q_ptr @@ -1136,7 +877,8 @@ def qk_norm_and_half_rope_backward_kernel( + i * L * stride + D + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) if SILU: @@ -1172,6 +914,8 @@ def qk_norm_and_half_rope_backward_kernel( 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 @@ -1179,6 +923,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dq_0, + mask=row_mask, ) tl.store( dq_ptr @@ -1188,8 +933,12 @@ def qk_norm_and_half_rope_backward_kernel( + 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 @@ -1197,6 +946,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dq_0, + mask=row_mask, ) tl.store( dq_ptr @@ -1206,6 +956,7 @@ def qk_norm_and_half_rope_backward_kernel( + 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) @@ -1217,11 +968,15 @@ def qk_norm_and_half_rope_backward_kernel( 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] @@ -1230,25 +985,27 @@ def qk_norm_and_half_rope_backward_kernel( gk_ptr + i * L * h * DD + pid * h * DD - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :] + + 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, :] + + 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, 2, d)), (0, 2, 1)), dim=2) + tl.flip(tl.permute(tl.reshape(gk_0, (h_p, 2, d)), (0, 2, 1)), dim=2) * signs, (0, 2, 1), ), - (h, D), + (h_p, D), ) gk_0 = gk_0 * cos + gk_r * sin @@ -1258,7 +1015,8 @@ def qk_norm_and_half_rope_backward_kernel( + pid * B * stride + i * stride + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) k1 = tl.load( k_ptr @@ -1266,7 +1024,8 @@ def qk_norm_and_half_rope_backward_kernel( + i * stride + D + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) else: k0 = tl.load( @@ -1274,7 +1033,8 @@ def qk_norm_and_half_rope_backward_kernel( + pid * stride + i * L * stride + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) k1 = tl.load( k_ptr @@ -1282,7 +1042,8 @@ def qk_norm_and_half_rope_backward_kernel( + i * L * stride + D + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) if SILU: @@ -1326,6 +1087,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dk_0, + mask=row_mask, ) tl.store( dk_ptr @@ -1335,6 +1097,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dk_1, + mask=row_mask, ) else: tl.store( @@ -1344,6 +1107,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dk_0, + mask=row_mask, ) tl.store( dk_ptr @@ -1353,17 +1117,23 @@ def qk_norm_and_half_rope_backward_kernel( + 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): @@ -1372,16 +1142,18 @@ def qk_norm_and_half_rope_backward_kernel( gv_ptr + i * L * h * DD + pid * h * DD - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :] + + 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)[:, None] - + tl.arange(0, D)[None, :] + + 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: @@ -1391,7 +1163,8 @@ def qk_norm_and_half_rope_backward_kernel( + pid * B * stride + i * stride + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) v1 = tl.load( v_ptr @@ -1399,7 +1172,8 @@ def qk_norm_and_half_rope_backward_kernel( + i * stride + D + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) else: v0 = tl.load( @@ -1407,7 +1181,8 @@ def qk_norm_and_half_rope_backward_kernel( + i * L * stride + pid * stride + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) v1 = tl.load( v_ptr @@ -1415,7 +1190,8 @@ def qk_norm_and_half_rope_backward_kernel( + pid * stride + D + DD * row_offs[:, None] - + tl.arange(0, D)[None, :] + + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) s0 = tl.sigmoid(v0) @@ -1434,6 +1210,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dv_0, + mask=row_mask, ) tl.store( dv_ptr @@ -1443,6 +1220,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dv_1, + mask=row_mask, ) else: tl.store( @@ -1452,6 +1230,7 @@ def qk_norm_and_half_rope_backward_kernel( + DD * row_offs[:, None] + tl.arange(0, D)[None, :], dv_0, + mask=row_mask, ) tl.store( dv_ptr @@ -1461,1280 +1240,272 @@ def qk_norm_and_half_rope_backward_kernel( + 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: + 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] + 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 + h = gk.shape[2] + stride = qkv.stride(1) + + dtype = gq.dtype + device = gq.device + if transposed: + dqkv = torch.empty((L, B, (H + 2 * h) * D), dtype=dtype, device=device) + else: + dqkv = torch.empty((B, L, (H + 2 * h) * D), dtype=dtype, device=device) + grad_stride = dqkv.stride(1) # for potential fused kernel + + 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, + qkv, + q_norm_weight, + k_norm_weight, + freqs, + dqkv, + 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 compatible_qk_norm_and_half_rope_backward_kernel( - gq_ptr, - gk_ptr, - gv_ptr, +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, - dqkv_ptr, - dqw_ptr, - dkw_ptr, - B, + cu_seqlens_q_ptr, + cu_seqlens_kv_ptr, + qo_ptr, + ko_ptr, + vo_ptr, stride, - grad_stride, eps, + mscale, + cp_rank, + B, + PB: tl.constexpr, H: tl.constexpr, h: tl.constexpr, - H_p: tl.constexpr, - h_p: tl.constexpr, + PH: tl.constexpr, + ph: tl.constexpr, D: tl.constexpr, d: tl.constexpr, INTERLEAVED: tl.constexpr, - TRANSPOSED: tl.constexpr, SILU: tl.constexpr, + CP_SIZE: tl.constexpr, + REUSE: 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)).to(tl.float32) - cos = tl.cos(freqs) - sin = tl.sin(freqs) - signs = -tl.arange(0, 2).to(tl.float32) * 2 + 1 + pos = _get_varlen_token_idx(cu_seqlens_q_ptr, pid, B, PB, cp_rank, CP_SIZE) - q_w0 = tl.load(q_norm_weight_ptr + tl.arange(0, D)) - q_w1 = tl.load(q_norm_weight_ptr + D + tl.arange(0, D)) + DD = D * 2 - dqw_0 = tl.zeros((D,), dtype=tl.float32) - dqw_1 = tl.zeros((D,), dtype=tl.float32) + 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 - dq_ptr = dqkv_ptr - # [bs, len, q_head, head_dim] -> [len, bs, q_head, head_dim] + w = H // h + + # [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_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_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_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_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, :]) - # 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 + pid * stride + i * L * stride + DD * row_offs[:,None] + tl.arange(0, D)[None, :]) - # q1 = tl.load(q_ptr + pid * stride + i * L * stride + D + 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, :], - 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] - - 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, - 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, - 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)).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_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_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_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_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, :], - 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, :], - 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, :], - 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, :], - 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] - - if TRANSPOSED: - 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, - 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_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_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): - - 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 - + 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 - + 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: - 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] - 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 - h = gk.shape[2] - stride = qkv.stride(1) - - dtype = gq.dtype - device = gq.device - if transposed: - dqkv = torch.empty((L, B, (H + 2 * h) * D), dtype=dtype, device=device) - else: - dqkv = torch.empty((B, L, (H + 2 * h) * D), dtype=dtype, device=device) - grad_stride = dqkv.stride(1) # for potential fused kernel - - 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,) - if H == H_p and h == h_p: - qk_norm_and_half_rope_backward_kernel[grid]( - gq, - gk, - gv, - qkv, - q_norm_weight, - k_norm_weight, - freqs, - dqkv, - tmp_dqw, - tmp_dkw, - B, - stride, - grad_stride, - eps, - H, - h, - D // 2, - D // 4, - interleaved, - transposed, - silu, - num_stages=num_stages, - num_warps=num_warps, - ) - - else: - compatible_qk_norm_and_half_rope_backward_kernel[grid]( - gq, - gk, - gv, - qkv, - q_norm_weight, - k_norm_weight, - freqs, - dqkv, - 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, 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, - 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, H) + tl.arange(0, H) // w * 2 - else: - row_offs = tl.arange(0, H) - - q0 = tl.load( - q_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).to(tl.float32) - q1 = tl.load( - q_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).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, H)[:, None] - + tl.arange(0, D)[None, :], - q1, - ) - - 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), - ) - q0 = q0 * cos + qr * sin - tl.store( - qo_ptr - + pid * 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)).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, h) * (w + 2) - k_ptr = qkv_ptr + DD * w - else: - row_offs = tl.arange(0, h) - k_ptr = qkv_ptr + DD * H - - k0 = tl.load( - k_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).to(tl.float32) - k1 = tl.load( - k_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).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 - tl.store( - ko_ptr - + pid * h * DD - + D - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :], - k1, - ) - - 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), - ) - k0 = k0 * cos + kr * sin - tl.store( - ko_ptr - + pid * h * DD - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :], - k0, - ) - - if INTERLEAVED: - row_offs = tl.arange(0, h) * (w + 2) - v_ptr = qkv_ptr + DD * w + DD - else: - row_offs = tl.arange(0, 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, :] - ).to(tl.float32) - v1 = tl.load( - v_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).to(tl.float32) - - if SILU: - v0 = v0 * tl.sigmoid(v0) - v1 = v1 * tl.sigmoid(v1) - - tl.store( - vo_ptr - + pid * h * DD - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :], - v0, - ) - tl.store( - vo_ptr - + pid * h * DD - + D - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :], - v1, - ) - - -@triton.jit -def compatible_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 + 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 = tl.arange(0, ph)[:, None] < h - v_ptr = qkv_ptr + DD * H + DD * h + row_offs = tl.arange(0, H) + row_mask = row_offs[:, None] < H + q_mask = tl.arange(0, PH)[:, None] < H - v0 = tl.load( - v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + q0 = tl.load( + q_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, :], + 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: - 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, - ) + 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( - vo_ptr - + pid * h * DD + qo_ptr + + pid * H * DD + D - + DD * tl.arange(0, ph)[:, None] + + DD * tl.arange(0, PH)[:, None] + tl.arange(0, D)[None, :], - v1, - mask=v_mask, + q1, + mask=q_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 - stride = qkv.stride(0) # qkv may be a slice of a tensor - D = Dim // (H + 2 * h) - B = cu_seqlens_q.size(0) - 1 - PB = max(triton.next_power_of_2(B), 128) # reduce jit - 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) - - if PH == H and ph == 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, - D // 2, - D // 4, - interleaved, - silu, - cp_size, - reuse, - num_stages=num_stages, - num_warps=num_warps, - ) - else: - compatible_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, - 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, H) + tl.arange(0, H) // w * 2 - else: - row_offs = tl.arange(0, H) - - gq_0 = tl.load( - gq_ptr + pid * H * DD + DD * tl.arange(0, H)[:, None] + tl.arange(0, D)[None, :] - ).to(tl.float32) - gq_1 = tl.load( - gq_ptr - + pid * H * DD - + D - + DD * tl.arange(0, H)[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - - gq_r = tl.reshape( + q0 *= rms[:, None] + q0 *= q_weight_0 + qr = tl.reshape( tl.permute( - tl.flip(tl.permute(tl.reshape(gq_0, (H, 2, d)), (0, 2, 1)), dim=2) * signs, + tl.flip(tl.permute(tl.reshape(q0, (PH, 2, d)), (0, 2, 1)), dim=2) * signs, (0, 2, 1), ), - (H, 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, :] - ).to(tl.float32) - q1 = tl.load( - q_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).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, + (PH, D), ) + q0 = q0 * cos + qr * sin tl.store( - dq_ptr - + pid * grad_stride - + D - + DD * row_offs[:, None] + qo_ptr + + pid * H * DD + + DD * tl.arange(0, PH)[:, None] + tl.arange(0, D)[None, :], - dq_1, + q0, + mask=q_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_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) @@ -2742,147 +1513,200 @@ def varlen_qk_norm_and_half_rope_backward_kernel( 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, h) * (w + 2) + row_offs = tl.arange(0, ph) * (w + 2) k_ptr = qkv_ptr + DD * w - dk_ptr = dqkv_ptr + DD * w + row_mask = row_offs[:, None] < (h * (w + 2)) else: - row_offs = tl.arange(0, h) + row_offs = tl.arange(0, ph) 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, h)[:, None] + tl.arange(0, D)[None, :] - ).to(tl.float32) - gk_1 = tl.load( - gk_ptr - + pid * h * DD - + D - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :] - ).to(tl.float32) - - 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_0 = gk_0 * cos + gk_r * sin + row_mask = tl.arange(0, ph)[:, None] < h k0 = tl.load( - k_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] + 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, :] + 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, - ) + 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( - dk_ptr - + pid * grad_stride + ko_ptr + + pid * h * DD + D - + DD * row_offs[:, None] + + DD * tl.arange(0, ph)[:, None] + tl.arange(0, D)[None, :], - dk_1, + k1, + mask=k_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) + 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, + ) - # [t, k_head, head_dim] -> [t, k_head + 2 * kv_head, head_dim] if INTERLEAVED: - row_offs = tl.arange(0, h) * (w + 2) + 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, h) + row_offs = tl.arange(0, ph) + row_mask = tl.arange(0, ph)[:, 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, h)[:, None] + tl.arange(0, D)[None, :] + v0 = tl.load( + v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], + mask=row_mask, ).to(tl.float32) - gv_1 = tl.load( - gv_ptr - + pid * h * DD - + D - + DD * tl.arange(0, h)[:, None] - + tl.arange(0, D)[None, :] + 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 = tl.load( - v_ptr + pid * stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).to(tl.float32) - v1 = tl.load( - v_ptr + pid * stride + D + DD * row_offs[:, None] + tl.arange(0, D)[None, :] - ).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 + v0 = v0 * tl.sigmoid(v0) + v1 = v1 * tl.sigmoid(v1) + v_mask = tl.arange(0, ph)[:, None] < h tl.store( - dv_ptr + pid * grad_stride + DD * row_offs[:, None] + tl.arange(0, D)[None, :], - dv_0, + vo_ptr + + pid * h * DD + + DD * tl.arange(0, ph)[:, None] + + tl.arange(0, D)[None, :], + v0, + mask=v_mask, ) tl.store( - dv_ptr - + pid * grad_stride + vo_ptr + + pid * h * DD + D - + DD * row_offs[:, None] + + DD * tl.arange(0, ph)[:, None] + tl.arange(0, D)[None, :], - dv_1, + 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 compatible_varlen_qk_norm_and_half_rope_backward_kernel( +def varlen_qk_norm_and_half_rope_backward_kernel( gq_ptr, gk_ptr, gv_ptr, @@ -3236,7 +2060,7 @@ def triton_varlen_qk_norm_and_half_rope_backward( stride = qkv.stride(0) h = gk.shape[1] B = cu_seqlens_q.size(0) - 1 - PB = max(triton.next_power_of_2(B), 128) + PB = 128 num_stages = 5 num_warps = 1 @@ -3253,72 +2077,39 @@ def triton_varlen_qk_norm_and_half_rope_backward( PH = triton.next_power_of_2(H) ph = triton.next_power_of_2(h) - if PH == H and ph == 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, - D // 2, - D // 4, - interleaved, - silu, - cp_size, - reuse, - num_stages=num_stages, - num_warps=num_warps, - ) - else: - compatible_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, - ) + 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, + ) dqw = tmp_dqw.sum(0) dkw = tmp_dkw.sum(0) return dqkv, dqw, dkw @@ -3541,7 +2332,7 @@ def triton_mla_rope_forward( assert cu_seqlens_kv is not None N, H, D = q.shape B = cu_seqlens_q.shape[0] - 1 - PB = max(triton.next_power_of_2(B), 128) + PB = 128 qo = None ko = torch.empty((N, H, 192), dtype=dtype, device=device) vo = torch.empty((N, H, 128), dtype=dtype, device=device) @@ -3808,8 +2599,7 @@ def triton_mla_rope_backward( assert cu_seqlens_kv is not None N, H, D = q_grad.shape B = cu_seqlens_q.shape[0] - 1 - PB = max(triton.next_power_of_2(B), 128) - assert B <= 128 + PB = 128 dq = None dkv = torch.empty((N, H, 256), dtype=dtype, device=device) dp = torch.empty((N, 1, 64), dtype=dtype, device=device) diff --git a/linghe/utils/scatter.py b/linghe/utils/scatter.py index 501c489..27a57ea 100644 --- a/linghe/utils/scatter.py +++ b/linghe/utils/scatter.py @@ -240,3 +240,91 @@ def triton_unpermute_with_mask_map( 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 b79f744..af58d5e 100644 --- a/linghe/utils/silu.py +++ b/linghe/utils/silu.py @@ -217,10 +217,12 @@ def silu_and_block_quant_forward_kernel( 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, ): @@ -236,11 +238,14 @@ def silu_and_block_quant_forward_kernel( 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) - x = x1 * tl.sigmoid(x1) * x2 - 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))) @@ -260,11 +265,19 @@ def silu_and_block_quant_forward_kernel( ) 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(transpose_output_ptr.dtype.element_ty) + xq = (x / scale).to(out_ptr.dtype.element_ty) tl.store( transpose_output_ptr + rid * H @@ -277,7 +290,7 @@ def silu_and_block_quant_forward_kernel( def triton_silu_and_block_quant_forward( - x, out=None, scale=None, round_scale=False, output_mode=2 + x, out=None, scale=None, limit=None, round_scale=False, output_mode=2 ): """ fused silu and blockwise quantization, used in shared expert @@ -303,17 +316,21 @@ def triton_silu_and_block_quant_forward( 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, 4 + H, W, num_warps = 64, 128, 8 elif output_mode == 1: - H, W, num_warps = 128, 64, 4 + 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]( @@ -322,13 +339,15 @@ def triton_silu_and_block_quant_forward( scale, transpose_output, transpose_scale, + limit, M, n, H, W, + CUTOFF, round_scale, output_mode, - num_stages=2, + num_stages=3, num_warps=num_warps, ) @@ -343,8 +362,10 @@ def silu_and_block_quant_backward_kernel( 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) @@ -366,6 +387,9 @@ def silu_and_block_quant_backward_kernel( 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 @@ -376,6 +400,9 @@ def silu_and_block_quant_backward_kernel( ).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))) @@ -396,7 +423,10 @@ def silu_and_block_quant_backward_kernel( qdx1 = (dx1 / scale1[None, :]).to(transpose_dx_ptr.dtype.element_ty) tl.store(transpose_dx_ptr + toffs, tl.trans(qdx1), mask=idx[None, :] < 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) if ROUND: scale2 = tl.exp2(tl.ceil(tl.log2(scale2))) @@ -421,7 +451,7 @@ def silu_and_block_quant_backward_kernel( # used in shared expert -def triton_silu_and_block_quant_backward(g, x, round_scale=False): +def triton_silu_and_block_quant_backward(g, x, limit=None, round_scale=False): """ backward of triton_silu_and_block_quant_forward Args: @@ -446,8 +476,14 @@ def triton_silu_and_block_quant_backward(g, x, round_scale=False): transpose_dx = torch.empty((N, M), device=device, dtype=torch.float8_e4m3fn) transpose_dx_scale = torch.empty(scale_shape, device=device, dtype=torch.float32) - assert M % 128 == 0 and N % 256 == 0 - grid = (M // 128, N // 256) + 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, @@ -455,8 +491,10 @@ def triton_silu_and_block_quant_backward(g, x, round_scale=False): dx_scale, transpose_dx, transpose_dx_scale, + limit, M, n, + CUTOFF, round_scale, num_stages=2, num_warps=8, @@ -474,8 +512,10 @@ def batch_weighted_silu_and_block_quant_forward_kernel( transpose_scale_ptr, count_ptr, accum_ptr, + limit, n, E: tl.constexpr, + CUTOFF: tl.constexpr, ROUND: tl.constexpr, ): eid = tl.program_id(axis=0) @@ -511,47 +551,53 @@ def batch_weighted_silu_and_block_quant_forward_kernel( + 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) 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) - 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] - scale = tl.maximum(tl.max(tl.abs(x), 1) / 448, 1e-30) + 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: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + 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), - scale, + scale1, mask=indices < count, ) + xq1 = (x / scale1[:, None]).to(out_ptr.dtype.element_ty) + tl.store(out_ptr + hoffs, xq1, mask=mask) - xq = (x / scale[:, None]).to(out_ptr.dtype.element_ty) - tl.store(out_ptr + hoffs, xq, mask=mask) - - 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 * 128 + tl.arange(0, 128), - scale, + scale2, ) - - xq = tl.trans((x / scale).to(out_ptr.dtype.element_ty)) - tl.store(transpose_output_ptr + toffs, xq, mask=indices[None, :] < count) + 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 @@ -564,9 +610,11 @@ def batch_weighted_silu_and_block_quant_forward_nt_kernel( 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) @@ -601,13 +649,22 @@ def batch_weighted_silu_and_block_quant_forward_nt_kernel( ) 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] - 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: @@ -622,45 +679,64 @@ def batch_weighted_silu_and_block_quant_forward_nt_kernel( soffs += B indices += B - # transpose - counts = tl.load(count_ptr + tl.arange(0, E)) - n_blocks = tl.cdiv(counts, 128) - transpose_soff = 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, B)[None, :] + + 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, B)[:, None] * count - + tl.arange(0, 128)[None, :] + + tl.arange(0, 128)[:, None] * count + + tl.arange(0, B)[None, :] ) - tsoffs = transpose_soff * n + rid * n + cid * 128 + tl.arange(0, B) - indices = rid * 128 + tl.arange(0, 128) 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] - 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 + tsoffs, scale) - - xq = tl.trans((x / scale).to(transpose_output_ptr.dtype.element_ty)) + xq = tl.trans((x / scale[None, :])) tl.store(transpose_output_ptr + toffs, xq, mask=indices[None, :] < count) - offs += B - toffs += count * B - tsoffs += B + offs += B * n * 2 + hoffs += B * n + indices += B + toffs += B @triton.jit @@ -671,9 +747,11 @@ def batch_weighted_silu_and_block_quant_forward_n_kernel( 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) @@ -703,7 +781,14 @@ def batch_weighted_silu_and_block_quant_forward_n_kernel( 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 * 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] scale = tl.maximum(tl.max(tl.abs(x), 1) / 448, 1e-30) if ROUND: @@ -733,9 +818,11 @@ def batch_weighted_silu_and_block_quant_forward_t_kernel( 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) @@ -768,7 +855,14 @@ def batch_weighted_silu_and_block_quant_forward_t_kernel( 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 * 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] scale = tl.maximum(tl.max(tl.abs(x), 0) / 448, 1e-30) if ROUND: @@ -801,6 +895,7 @@ def triton_batch_weighted_silu_and_block_quant_forward( splits=None, out=None, scale=None, + limit=None, round_scale=False, output_mode=2, ): @@ -842,6 +937,12 @@ def triton_batch_weighted_silu_and_block_quant_forward( 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: @@ -854,9 +955,11 @@ def triton_batch_weighted_silu_and_block_quant_forward( scale, counts, accums, + limit, n, B, len(splits), + CUTOFF, round_scale, num_stages=2, num_warps=2, @@ -871,16 +974,19 @@ def triton_batch_weighted_silu_and_block_quant_forward( 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_kernel[grid]( + batch_weighted_silu_and_block_quant_forward_nt_kernel[grid]( x, weight, out, @@ -889,64 +995,15 @@ def triton_batch_weighted_silu_and_block_quant_forward( transpose_scale, counts, accums, + limit, n, + B, len(splits), + CUTOFF, round_scale, - num_stages=2, - num_warps=8, + num_stages=3, + num_warps=4, ) - - # 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, - # n, - # B, - # len(splits), - # round_scale, - # num_stages=2, - # num_warps=8 - # ) - - # 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, - # n, - # B, - # len(splits), - # round_scale, - # num_stages=2, - # num_warps=2 - # ) - # 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, - # n, - # B, - # len(splits), - # round_scale, - # num_stages=2, - # num_warps=2 - # ) return out, scale, transpose_output, transpose_scale @@ -962,8 +1019,10 @@ def batch_weighted_silu_and_block_quant_backward_kernel( 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) @@ -1001,6 +1060,15 @@ def batch_weighted_silu_and_block_quant_backward_kernel( 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 @@ -1010,36 +1078,38 @@ def batch_weighted_silu_and_block_quant_backward_kernel( + tl.arange(0, 128)[None, :], mask=idx[:, None] < count, ).to(tl.float32) - sigmoid = tl.sigmoid(x1) - dw = tl.sum(sigmoid * x1 * x2 * g, 1) + 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 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: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + 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), - scale, + scale1, mask=idx < count, ) - tl.store(dx_ptr + offs, dx / scale[:, None], mask=idx[:, None] < count) + tl.store(dx_ptr + offs, dx / scale1[:, 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))) tl.store( transpose_dx_scale_ptr + transpose_off * n * 2 + rid * n * 2 + cid * 128 + tl.arange(0, 128), - scale, + scale2, ) - qdx = tl.trans((dx / scale[None, :]).to(dx_ptr.dtype.element_ty)) + 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 @@ -1052,10 +1122,15 @@ def batch_weighted_silu_and_block_quant_backward_kernel( mask=idx[None, :] < count, ) - dx = sigmoid * g * x1 * w - scale = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + dx = g * gate * w + + if CUTOFF: + dx *= dx2_mask + + scale3 = tl.maximum(tl.max(dx.abs(), 1) / 448, 1e-30) + if ROUND: - scale = tl.exp2(tl.ceil(tl.log2(scale))) + scale3 = tl.exp2(tl.ceil(tl.log2(scale3))) tl.store( dx_scale_ptr + si * nb * 2 @@ -1063,15 +1138,15 @@ def batch_weighted_silu_and_block_quant_backward_kernel( + rid * 128 + count * nb + tl.arange(0, 128), - scale, + scale3, mask=idx < count, ) - tl.store(dx_ptr + n + offs, dx / scale[:, None], mask=idx[:, None] < count) + tl.store(dx_ptr + n + offs, dx / scale3[:, None], mask=idx[:, None] < count) - scale = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) + scale4 = 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)) + 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 @@ -1079,7 +1154,7 @@ def batch_weighted_silu_and_block_quant_backward_kernel( + n + cid * 128 + tl.arange(0, 128), - scale, + scale4, ) tl.store( transpose_dx_ptr @@ -1104,9 +1179,11 @@ def batch_weighted_silu_and_block_quant_backward_n_kernel( 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) @@ -1136,6 +1213,15 @@ def batch_weighted_silu_and_block_quant_backward_n_kernel( 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 @@ -1145,12 +1231,15 @@ def batch_weighted_silu_and_block_quant_backward_n_kernel( + tl.arange(0, 128)[None, :], mask=idx[:, None] < count, ).to(tl.float32) - sigmoid = tl.sigmoid(x1) - dw = tl.sum(sigmoid * x1 * x2 * g, 1) + 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))) @@ -1162,7 +1251,11 @@ def batch_weighted_silu_and_block_quant_backward_n_kernel( tl.store(dx_ptr + offs, dx / scale[:, None], mask=idx[:, None] < count) - dx = sigmoid * g * x1 * w + 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))) @@ -1188,9 +1281,11 @@ def batch_weighted_silu_and_block_quant_backward_t_kernel( 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) @@ -1224,6 +1319,16 @@ def batch_weighted_silu_and_block_quant_backward_t_kernel( 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 @@ -1233,10 +1338,12 @@ def batch_weighted_silu_and_block_quant_backward_t_kernel( + tl.arange(0, B)[None, :], mask=idx[:, None] < count, ).to(tl.float32) - sigmoid = tl.sigmoid(x1) 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))) @@ -1261,7 +1368,10 @@ def batch_weighted_silu_and_block_quant_backward_t_kernel( mask=idx[None, :] < count, ) - dx = sigmoid * g * x1 * w + dx = g * gate * w + + if CUTOFF: + dx *= dx2_mask scale = tl.maximum(tl.max(dx.abs(), 0) / 448, 1e-30) if ROUND: @@ -1291,7 +1401,7 @@ def batch_weighted_silu_and_block_quant_backward_t_kernel( # used in routed experts def triton_batch_weighted_silu_and_block_quant_backward( - g, x, weight, counts, splits=None, round_scale=False + g, x, weight, counts, splits=None, limit=None, round_scale=False ): """ backward of triton_batch_weighted_silu_and_block_quant_forward @@ -1330,6 +1440,12 @@ def triton_batch_weighted_silu_and_block_quant_backward( 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]( @@ -1343,8 +1459,10 @@ def triton_batch_weighted_silu_and_block_quant_backward( # transpose_dx, # transpose_dx_scale, # dws, + # limit, # n, # n_expert, + # CUTOFF, # round_scale, # num_stages=2, # num_warps=8 @@ -1363,9 +1481,11 @@ def triton_batch_weighted_silu_and_block_quant_backward( dx, dx_scale, dws, + limit, n, B, n_expert, + CUTOFF, round_scale, num_stages=2, num_warps=4, @@ -1382,9 +1502,11 @@ def triton_batch_weighted_silu_and_block_quant_backward( accums, transpose_dx, transpose_dx_scale, + limit, n, B, n_expert, + CUTOFF, round_scale, num_stages=2, num_warps=4, @@ -1393,84 +1515,758 @@ def triton_batch_weighted_silu_and_block_quant_backward( return dx, dx_scale, dw, transpose_dx, transpose_dx_scale -# n is power of 2 @triton.jit -def silu_and_smooth_quant_forward_kernel( +def silu_and_mxfp8_quant_forward_kernel( x_ptr, - smooth_scale_ptr, out_ptr, scale_ptr, - max_ptr, + transpose_output_ptr, + transpose_scale_ptr, + limit, M, - T, + m, n: tl.constexpr, - W: tl.constexpr, - ROUND: tl.constexpr, - CALIBRATE: tl.constexpr, + B: tl.constexpr, + CUTOFF: tl.constexpr, + OUTPUT_MODE: 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) + rid = tl.program_id(axis=0) + cid = tl.program_id(axis=1) + 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 - 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) + 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 - if CALIBRATE: - maxs = tl.maximum(x.abs(), maxs) - x = x * smooth_scale - 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 + indices, scale, mask=indices < M) - x = (x / scale[:, None]).to(out_ptr.dtype.element_ty) - 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) + 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( + 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, + ) -# 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, + 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( + 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 ): - pid = tl.program_id(axis=0) + """ + fused silu and mxfp8 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 - # rowwise read with block size [T, B] - row_offs = pid * T * n + tl.arange(0, T)[:, None] * n - col_offs = tl.arange(0, B)[None, :] + 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 + 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), device=device, dtype=torch.float8_e4m3fn) + if scale is None: + scale = torch.empty((M, n // 32), device=device, dtype=torch.uint8) - nb = n // B - maxs = tl.zeros((T,), dtype=tl.float32) - for i in range(nb): + if limit is None: + CUTOFF = False + limit = 0.0 + else: + CUTOFF = True + 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, + B, + CUTOFF, + output_mode, + num_stages=2, + num_warps=1, + ) + + return out, scale, transpose_output, transpose_scale + + +@triton.jit +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 // 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 + 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) + + log_scale1 = tl.ceil(tl.log2(scale1)) + scale1 = tl.exp2(log_scale1) + + tl.store( + transpose_dx_scale_ptr + rid * n * 2 + cid * 32 + tl.arange(0, 32), + log_scale1 + 127, + ) + + qdx1 = (dx1 / scale1[None, :]).to(transpose_dx_ptr.dtype.element_ty) + tl.store(transpose_dx_ptr + offs, qdx1, mask=mask) + + # 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=mask) + + 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(transpose_dx_ptr.dtype.element_ty) + tl.store(transpose_dx_ptr + n + offs, qdx2, mask=mask) + + +# used in shared expert +def triton_silu_and_mxfp8_quant_backward(g, x, limit=None): + """ + backward of triton_silu_and_mxfp8_quant_forward + Args: + g: gradient + x: input tensor + + Returns: + - dx: rowwise quantized gradient + - dx_scale: scales of rowwise quantized gradient + - transpose_dx: columnwise quantized gradient + - transpose_dx_scale: scales of columnwise quantized gradient + """ + 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((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 + + 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, + CUTOFF, + num_stages=3, + num_warps=4, + ) + return dx, dx_scale, transpose_dx, transpose_dx_scale + + +@triton.jit +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) + + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) + si = ei - count + c = tl.cdiv(count, 128) + + if rid >= c * 4: + return + + n = n.to(tl.int64) + nb = n // 32 + + counts = tl.load(count_ptr + tl.arange(0, E)) + n_blocks = tl.cdiv(counts, 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) + 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] + + if OUTPUT_MODE % 2 == 0: + scale = tl.maximum(tl.max(tl.abs(x), 1) / 448, 1e-30) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + # 4 = 128 // 32 + tl.store( + 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) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store( + 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_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: + 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 + 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() and weight.is_contiguous() + m, N = x.shape + n = N // 2 + 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 + if out is None: + out = torch.empty((m, n), device=device, dtype=torch.float8_e4m3fn) + + if scale is None: + 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) * 4, n // 32) + batch_weighted_silu_and_mxfp8_quant_forward_kernel[grid]( + x, + weight, + out, + scale, + transpose_output, + transpose_scale, + counts, + accums, + limit, + n, + len(splits), + CUTOFF, + output_mode, + num_stages=3, + num_warps=1, + ) + + return out, scale, transpose_output, transpose_scale + + +@triton.jit +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) + + count = tl.load(count_ptr + eid) + si = tl.load(accum_ptr + eid) - count + + if rid >= tl.cdiv(count, 128) * 4: + return + + 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 * 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 * 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) + 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)) + + 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) + log_scale = tl.ceil(tl.log2(scale)) + scale = tl.exp2(log_scale) + tl.store( + transpose_dx_scale_ptr + + scale_off * n * 8 + + rid * n * 2 + + cid * 32 + + tl.arange(0, 32), + log_scale + 127, + ) + + 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 * 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 + + 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 + + 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) + 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 + + scale_off * n * 8 + + rid * n * 2 + + n + + cid * 32 + + tl.arange(0, 32), + log_scale + 127, + ) + tl.store( + 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_mxfp8_quant_backward( + g, x, weight, counts, splits=None, limit=None +): + """ + 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 + 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 + """ + assert g.is_contiguous() + m, N = x.shape + n = N // 2 + 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) + 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 M == 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_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, + counts, + accums, + dx, + dx_scale, + transpose_dx, + transpose_dx_scale, + dws, + limit, + n, + n_experts, + CUTOFF, + num_stages=3, + 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, + 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 + + 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 * tl.sigmoid(x1) * x2 + x = x * smooth_scale + 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 + indices, scale, mask=indices < M) + x = (x / scale[:, None]).to(out_ptr.dtype.element_ty) + tl.store(out_ptr + row_offs + col_offs, x, mask=mask) + row_offs += n * W + + +# 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, + 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] + row_offs = pid * T * n + tl.arange(0, T)[:, None] * n + col_offs = tl.arange(0, B)[None, :] + + 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 * tl.sigmoid(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 = x / smooth_scale maxs = tl.maximum(tl.max(x.abs(), 1), maxs) col_offs += B @@ -1496,13 +2292,7 @@ def compatible_silu_and_smooth_quant_forward_kernel( # 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, + x, smooth_scale, out=None, scale=None, round_scale=False ): """""" assert x.is_contiguous() @@ -1520,21 +2310,17 @@ def triton_silu_and_smooth_quant_forward( 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, ) @@ -1543,28 +2329,21 @@ def triton_silu_and_smooth_quant_forward( 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, ) - if calibrate: - maxs = maxs.amax(0) - - return out, scale, maxs + return out, scale @triton.jit @@ -1761,7 +2540,6 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel( smooth_scale_ptr, out_ptr, scale_ptr, - max_ptr, count_ptr, accum_ptr, M, @@ -1769,7 +2547,6 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel( W: tl.constexpr, ROUND: tl.constexpr, REVERSE: tl.constexpr, - CALIBRATE: tl.constexpr, ): eid = tl.program_id(axis=0) tid = tl.program_id(axis=1) @@ -1786,9 +2563,6 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel( 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 @@ -1800,9 +2574,6 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel( ] x = x1 * tl.sigmoid(x1) * x2 - if CALIBRATE: - maxs = tl.maximum(x.abs(), maxs) - x *= w * smooth_scale scale = tl.maximum(tl.max(x.abs(), 1) / 448, 1e-30) if ROUND: @@ -1812,10 +2583,6 @@ def batch_weighted_silu_and_smooth_quant_forward_kernel( 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( @@ -1828,7 +2595,6 @@ def triton_batch_weighted_silu_and_smooth_quant_forward( scale=None, round_scale=False, reverse=False, - calibrate=False, ): """""" assert x.is_contiguous() and weight.is_contiguous() @@ -1841,20 +2607,11 @@ def triton_batch_weighted_silu_and_smooth_quant_forward( out = torch.empty((M, n), device=device, dtype=torch.float8_e4m3fn) sm = 128 - tmp_maxs = None 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 @@ -1865,7 +2622,6 @@ def triton_batch_weighted_silu_and_smooth_quant_forward( smooth_scale, out, scale, - tmp_maxs, counts, accums, M, @@ -1873,14 +2629,11 @@ def triton_batch_weighted_silu_and_smooth_quant_forward( W, round_scale, reverse, - calibrate, num_stages=3, num_warps=16, ) - if calibrate: - maxs = tmp_maxs.amax(1) - return out, scale, maxs + return out, scale @triton.jit @@ -1915,16 +2668,8 @@ def batch_weighted_silu_and_smooth_quant_backward_kernel( 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 - ) + 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 @@ -2096,13 +2841,8 @@ def _batch_requant_kernel( 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, - ) - ) + 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 diff --git a/linghe/utils/topk.py b/linghe/utils/topk.py index 1c7e782..1c32262 100644 --- a/linghe/utils/topk.py +++ b/linghe/utils/topk.py @@ -14,7 +14,7 @@ def topk_forward_kernel( ): pid = tl.program_id(axis=0) - xo = tl.load(input_ptr + pid * N + tl.arange(0, N)) + xo = tl.load(input_ptr + pid * N + tl.arange(0, N)).to(tl.float32) x = xo for i in range(K): @@ -34,7 +34,61 @@ def topk_forward_kernel( y = tl.where(y == val, -2e38, y) -def triton_topk_forward(x, k, dim=-1): +@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: @@ -46,20 +100,30 @@ def triton_topk_forward(x, k, dim=-1): """ device = x.device shape = x.shape - assert dim == -1 and len(shape) <= 3 - assert x.is_contiguous() - if len(shape) == 3: - M, B, N = shape - g = M * B - values = torch.empty((M, B, k), device=device, dtype=x.dtype) - indices = torch.empty((M, B, k), device=device, dtype=torch.int64) - else: - M, N = shape - g = M - values = torch.empty((M, k), device=device, dtype=x.dtype) - indices = torch.empty((M, k), device=device, dtype=torch.int64) + 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,) - topk_forward_kernel[grid](x, values, indices, N, k, num_stages=2, num_warps=2) + + 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 @@ -84,16 +148,11 @@ def triton_topk_backward(grad_output, indices, N, dim=-1): """ device = grad_output.device shape = grad_output.shape - assert dim == -1 and len(shape) <= 3 + assert dim == -1 assert grad_output.is_contiguous() - if len(shape) == 3: - M, B, k = shape - g = M * B - dx = torch.zeros((M, B, N), device=device, dtype=grad_output.dtype) - else: - M, k = shape - g = M - dx = torch.zeros((M, N), device=device, dtype=grad_output.dtype) + 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 @@ -102,7 +161,7 @@ def triton_topk_backward(grad_output, indices, N, dim=-1): @triton.jit -def group_topk_score_forward_kernel( +def deprecated_group_topk_score_forward_kernel( input_ptr, bias_ptr, prob_ptr, @@ -173,6 +232,59 @@ def group_topk_score_forward_kernel( 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, @@ -209,21 +321,38 @@ def triton_group_topk_score_forward( routing_map = torch.empty((M, N), device=device, dtype=torch.bool) BIAS = expert_bias is not None grid = (g,) - group_topk_score_forward_kernel[grid]( - x, - expert_bias, - probs, - routing_map, - scaling_factor, - eps, - N, - k, - num_groups, - group_topk, - BIAS, - num_stages=1, - num_warps=1, - ) + 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) diff --git a/linghe/utils/transpose.py b/linghe/utils/transpose.py index ae65cb6..81ae4e9 100644 --- a/linghe/utils/transpose.py +++ b/linghe/utils/transpose.py @@ -12,6 +12,7 @@ from linghe.tools.util import round_up + # os.environ["TRITON_PRINT_AUTOTUNING"] = "1" @@ -31,19 +32,14 @@ def transpose_kernel( 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), - ) + 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), + 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 @@ -83,19 +79,14 @@ def transpose_outer_dims_kernel( 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), - ) + 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), + 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): @@ -151,7 +142,6 @@ def triton_transpose(x: torch.Tensor, inner=True): num_warps=num_warps, ) else: - 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) @@ -173,7 +163,7 @@ def triton_transpose(x: torch.Tensor, inner=True): @triton.jit -def transpose_and_pad_kernel( +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) @@ -196,7 +186,7 @@ def transpose_and_pad_kernel( 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 @@ -211,7 +201,7 @@ def triton_transpose_and_pad(x, out=None, pad=True): # 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) @@ -223,7 +213,7 @@ 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]( + pad_transpose_kernel[grid]( x, out, M, N, P, H, W, EVEN, num_stages=num_stages, num_warps=num_warps ) return out @@ -279,7 +269,7 @@ def triton_batch_transpose(xs, xts=None): @triton.jit -def batch_transpose_and_pad_kernel( +def batch_pad_transpose_kernel( x_ptr, t_ptr, count_ptr, @@ -311,7 +301,7 @@ def batch_transpose_and_pad_kernel( 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: @@ -324,12 +314,11 @@ def triton_batch_transpose_and_pad(x, count_list, x_t=None, pad=True): x_t: output tensor """ assert x.is_contiguous() - assert pad # 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)), @@ -345,7 +334,7 @@ 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]( + batch_pad_transpose_kernel[grid]( x, x_t, counts, diff --git a/linghe/utils/unary.py b/linghe/utils/unary.py index eb9ce22..fbd194b 100644 --- a/linghe/utils/unary.py +++ b/linghe/utils/unary.py @@ -66,21 +66,55 @@ def triton_calculate_smooth_scale( return output +@triton.jit +def clip_kernel(x_ptr, clip_value, N, B: tl.constexpr, EVEN: tl.constexpr): + pid = tl.program_id(axis=0).to(tl.int64) + offs = pid * B + tl.arange(0, B) + if EVEN: + x = tl.load(x_ptr + offs) + xc = tl.minimum(tl.maximum(x, -clip_value), clip_value) + tl.store(x_ptr + offs, xc, mask=(tl.abs(x) > 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) - bid = tl.program_id(axis=1) + 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)) - else: + 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.to(tl.int64) * t * B + tl.arange(0, B) + 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) @@ -101,7 +135,7 @@ def triton_batch_clip(xs, clip_value=100.0): if len(xs) == 0: return dtype = xs[0].dtype - assert dtype in (torch.float32, torch.bfloat16) + 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 @@ -111,8 +145,12 @@ def triton_batch_clip(xs, clip_value=100.0): 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 + if dtype == torch.float32: + DT = 0 + elif dtype == torch.bfloat16: + DT = 1 + else: + DT = 2 T = 256 tensor_count = len(xs) B = 512 From 5cabff43d65808382320c65161681d77f74fc4b1 Mon Sep 17 00:00:00 2001 From: "liangchen.liangche" Date: Mon, 27 Apr 2026 11:38:22 +0800 Subject: [PATCH 09/11] sync ops add infer ops WIP: update test 1 WIP: update test 2 --- linghe/attn/mla.py | 30 +- linghe/experimental/gemm.py | 113 +++ linghe/facade/fp32_gemm.py | 57 +- linghe/facade/fp8_gemm.py | 7 + linghe/facade/gate.py | 79 ++- linghe/facade/gemm.py | 220 ++++++ linghe/facade/linear.py | 343 +++++++++ linghe/facade/loss.py | 24 +- linghe/facade/norm.py | 114 +++ linghe/facade/permutation.py | 900 ++++++++++++++++++++++- linghe/facade/quantization.py | 359 ++++++++++ linghe/facade/silu.py | 363 +++++++++- linghe/facade/topk.py | 10 +- linghe/gemm/blockwise_fp8_gemm.py | 100 +-- linghe/gemm/fp32_gemm.py | 189 ++--- linghe/gemm/mxfp8_gemm.py | 1066 ++++++++++++++++++++++++++++ linghe/infer/__init__.py | 0 linghe/infer/gate.py | 156 ++++ linghe/infer/gemm.py | 226 ++++++ linghe/infer/grouped_gemm.py | 146 ++++ linghe/infer/norm.py | 220 ++++++ linghe/infer/quant.py | 62 ++ linghe/infer/rope.py | 574 +++++++++++++++ linghe/infer/silu.py | 223 ++++++ linghe/infer/topk.py | 121 ++++ linghe/quant/block.py | 26 +- linghe/quant/group.py | 4 +- linghe/quant/mxfp8.py | 233 ++++++ linghe/quant/smooth.py | 862 +++++++--------------- linghe/tools/benchmark.py | 3 +- linghe/tools/check.py | 7 +- linghe/tools/util.py | 124 +++- tests/test_add.py | 121 ++-- tests/test_blockwise_fp8_gemm.py | 88 +-- tests/test_blockwise_quant.py | 55 +- tests/test_channel_quant.py | 58 +- tests/test_channelwise_fp8_gemm.py | 179 ++--- tests/test_dist_loss.py | 4 +- tests/test_embedding.py | 251 ++++--- tests/test_fp32_gemm.py | 340 +++++---- tests/test_gate.py | 303 +++++--- tests/test_gather.py | 745 +++++++++++-------- tests/test_group_quant.py | 29 +- tests/test_la.py | 115 +-- tests/test_loss.py | 182 ++--- tests/test_mla.py | 424 ++++------- tests/test_mul.py | 57 +- tests/test_rearange.py | 57 +- tests/test_reduce.py | 125 ++-- tests/test_rope.py | 569 +++++---------- tests/test_scatter.py | 42 +- tests/test_silu.py | 840 ++++++++++++++-------- tests/test_smooth_quant.py | 330 ++++++--- tests/test_unary.py | 144 ++-- 54 files changed, 8792 insertions(+), 3227 deletions(-) create mode 100644 linghe/experimental/gemm.py create mode 100644 linghe/facade/fp8_gemm.py create mode 100644 linghe/facade/gemm.py create mode 100644 linghe/facade/linear.py create mode 100644 linghe/facade/quantization.py create mode 100644 linghe/gemm/mxfp8_gemm.py create mode 100644 linghe/infer/__init__.py create mode 100644 linghe/infer/gate.py create mode 100644 linghe/infer/gemm.py create mode 100644 linghe/infer/grouped_gemm.py create mode 100644 linghe/infer/norm.py create mode 100644 linghe/infer/quant.py create mode 100644 linghe/infer/rope.py create mode 100644 linghe/infer/silu.py create mode 100644 linghe/infer/topk.py create mode 100644 linghe/quant/mxfp8.py diff --git a/linghe/attn/mla.py b/linghe/attn/mla.py index 18382dc..dce5eaf 100644 --- a/linghe/attn/mla.py +++ b/linghe/attn/mla.py @@ -260,7 +260,7 @@ def triton_mla_forward(q, k, v, causal=True, safe=True, clip_value=None): max_logits = torch.empty((B, H, L), dtype=torch.float32, device=q.device) softmax_scale = 128 ** (-0.5) - clip = clip_value is not None + 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: @@ -838,7 +838,7 @@ def triton_mla_backward( ds = torch.empty((B, H, L), dtype=torch.float32, device=device) softmax_scale = 128 ** (-0.5) - clip = clip_value is not None + clip = clip_value is not None and clip_value > 0.0 clip_value = clip_value * softmax_scale if clip else 0.0 M = 64 @@ -1090,7 +1090,7 @@ def triton_varlen_mla_forward( 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 + 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 @@ -1379,9 +1379,23 @@ def varlen_mla_rs_kernel( 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) + # 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 @@ -1437,7 +1451,7 @@ def triton_varlen_mla_backward( ds = torch.empty((H, T), dtype=torch.float32, device=device) softmax_scale = 128 ** (-0.5) - clip = clip_value is not None + clip = clip_value is not None and clip_value > 0.0 clip_value = clip_value * softmax_scale if clip else 0.0 M = 64 @@ -1506,7 +1520,7 @@ def triton_varlen_mla_backward( grid = (T, NB) num_warps = 4 num_stages = 3 - PB = max(triton.next_power_of_2(B), 128) + PB = 128 varlen_mla_rs_kernel[grid]( gq, qo, 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/facade/fp32_gemm.py b/linghe/facade/fp32_gemm.py index 462bf26..bcd6237 100644 --- a/linghe/facade/fp32_gemm.py +++ b/linghe/facade/fp32_gemm.py @@ -3,58 +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 - if len(shape) == 3: - input = input.view(shape[0] * shape[1], shape[2]) - logits = triton_fp32_gemm(input, weight) - - ctx.input_requires_grad = input.requires_grad - ctx.weight_requires_grad = weight.requires_grad - ctx.shape = shape - 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 - - dx = triton_fp32_gemm_for_backward(grad_output, weight) - if len(grad_shape) == 3: - 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) +# 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 index 66a37c3..47cbaed 100644 --- a/linghe/facade/gate.py +++ b/linghe/facade/gate.py @@ -8,6 +8,8 @@ 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, ) @@ -16,9 +18,11 @@ 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 @@ -30,7 +34,7 @@ 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 + dy, attn_output, gate, weight, eps=ctx.eps, group_size=ctx.group_size ) return dx, dg, dw, None, None @@ -42,6 +46,7 @@ def group_rms_norm_gate( 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) @@ -51,7 +56,79 @@ def group_rms_norm_gate( 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=input.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=input.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=ctx.input_requires_grad, + ) + + return dx, dg_out, dw, None, None diff --git a/linghe/facade/gemm.py b/linghe/facade/gemm.py new file mode 100644 index 0000000..91fcd4d --- /dev/null +++ b/linghe/facade/gemm.py @@ -0,0 +1,220 @@ +# -*- 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 + 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/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 8f5ff29..b588871 100644 --- a/linghe/facade/loss.py +++ b/linghe/facade/loss.py @@ -25,7 +25,7 @@ def forward(ctx, logits, labels, ignore_index=-100, inplace=False, tp_group=None 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, labels, tp_group, ignore_index=ignore_index + logits_view, labels, tp_group, ignore_index=ignore_index ) else: loss, sum_exp, max_logit = triton_softmax_cross_entropy_forward( @@ -35,6 +35,7 @@ def forward(ctx, logits, labels, ignore_index=-100, inplace=False, tp_group=None 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]) @@ -97,25 +98,30 @@ def softmax_cross_entropy( ) -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): """""" diff --git a/linghe/facade/norm.py b/linghe/facade/norm.py index 6ec651d..58fdf40 100644 --- a/linghe/facade/norm.py +++ b/linghe/facade/norm.py @@ -9,6 +9,8 @@ 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, ) @@ -110,3 +112,115 @@ def block_rms_norm(input, weight, rms, quantizer, cls, eps=1e-6, 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 + ) + ) + + # 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 + + @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 index 8178c04..27cdadb 100644 --- a/linghe/facade/permutation.py +++ b/linghe/facade/permutation.py @@ -1,14 +1,29 @@ +# -*- 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, ) -from linghe.utils.scatter import triton_unpermute_with_mask_map class _PaddedPermute(torch.autograd.Function): @@ -20,12 +35,15 @@ def forward( 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=16) - num_out_tokens = sum([(x + 15) // 16 * 16 for x in tokens_per_expert_list]) + 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 @@ -57,6 +75,7 @@ def backward(ctx, grad_output, grad_prob, grad_map): None, None, None, + None, ) @@ -66,6 +85,7 @@ def padded_permute( 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. @@ -85,6 +105,7 @@ def padded_permute( routing_map, tokens_per_expert_cuda_tensor, tokens_per_expert_list, + multiple, ) return permuted_input, permuted_probs, row_id_map @@ -324,3 +345,876 @@ def block_padded_unpermute( 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/silu.py b/linghe/facade/silu.py index fa0f377..3f4fe91 100644 --- a/linghe/facade/silu.py +++ b/linghe/facade/silu.py @@ -5,12 +5,20 @@ 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): + 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]) @@ -18,10 +26,12 @@ def forward(ctx, input, quantizer, grad_quantizer, cls): 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=quantizer.force_pow_2_scales + 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]) @@ -46,8 +56,9 @@ def backward(ctx, grad_output): (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=grad_quantizer.force_pow_2_scales + grad_output_view, input_view, round_scale=round_scale, limit=ctx.limit ) output = ctx.cls( shape=ctx.shape, @@ -62,11 +73,11 @@ def backward(ctx, grad_output): is_2D_scaled=False, ) - return output, None, None, None + return output, None, None, None, None -def block_silu_impl(input, quantizer, grad_quantizer, cls): - output = BlockSiluFunction.apply(input, quantizer, grad_quantizer, cls) +def block_silu_impl(input, quantizer, grad_quantizer, cls, limit=None): + output = BlockSiluFunction.apply(input, quantizer, grad_quantizer, cls, limit) return output @@ -81,6 +92,7 @@ def forward( quantizers, grad_quantizers, cls, + limit, is_recomputing, ): shape = input.shape @@ -89,6 +101,7 @@ def forward( ctx.shape = shape ctx.splits = splits ctx.cls = cls + ctx.limit = limit ctx.save_for_backward(input, weights, counts) if is_recomputing is None: @@ -98,13 +111,16 @@ def forward( else: output_mode = 0 - x_q, x_scale, xt_q, xt_scale = ( + 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, - round_scale=quantizers[0].force_pow_2_scales, + limit=limit, + round_scale=round_scale, output_mode=output_mode, ) ) @@ -127,14 +143,17 @@ def forward( def backward(ctx, grad_output): input, weights, counts = ctx.saved_tensors grad_quantizers = ctx.grad_quantizers - x_q, x_scale, wgrad, xt_q, xt_scale = ( + 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=grad_quantizers[0].force_pow_2_scales, + round_scale=round_scale, + limit=ctx.limit, ) ) output = ctx.cls( @@ -150,7 +169,7 @@ def backward(ctx, grad_output): is_2D_scaled=False, ) - return output, wgrad, None, None, None, None, None, None + return output, wgrad, None, None, None, None, None, None, None def block_batch_weighted_silu_impl( @@ -161,10 +180,330 @@ def block_batch_weighted_silu_impl( quantizers, grad_quantizers, cls, + limit=None, is_recomputing=None, ): assert input.ndim == 2 output = BlockBatchWeightedSiluFunction.apply( - input, weights, counts, splits, quantizers, grad_quantizers, cls, is_recomputing + 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/topk.py b/linghe/facade/topk.py index ed56d11..c41a520 100644 --- a/linghe/facade/topk.py +++ b/linghe/facade/topk.py @@ -17,8 +17,8 @@ class TopkFunction(torch.autograd.Function): """""" @staticmethod - def forward(ctx, x, k, dim): - values, indices = triton_topk_forward(x, k, dim=dim) + 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) @@ -30,10 +30,10 @@ def backward(ctx, grad_output, grad_indices): grad_input = triton_topk_backward( grad_output, indices, ctx.shape[-1], dim=ctx.dim ) - return grad_input, None, None + return grad_input, None, None, None, None -def fused_topk(x, k, dim=-1): +def fused_topk(x, k, dim=-1, sorted=True, impl="iter"): """ topk Args: @@ -44,7 +44,7 @@ def fused_topk(x, k, dim=-1): values: topk values indices: topk indices """ - return TopkFunction.apply(x, k, dim) + return TopkFunction.apply(x, k, dim, sorted, impl) class GroupTopkScoreFunction(torch.autograd.Function): diff --git a/linghe/gemm/blockwise_fp8_gemm.py b/linghe/gemm/blockwise_fp8_gemm.py index d2e022a..b2c0903 100644 --- a/linghe/gemm/blockwise_fp8_gemm.py +++ b/linghe/gemm/blockwise_fp8_gemm.py @@ -7,12 +7,12 @@ import triton import triton.language as tl -# adapt from deepseek + # os.environ["TRITON_PRINT_AUTOTUNING"] = "1" @triton.jit -def fp8_gemm_bb_kernel( +def fp8_blockwise_gemm_kernel( a_ptr, b_ptr, c_ptr, @@ -55,8 +55,8 @@ 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( +# 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, @@ -75,7 +75,7 @@ def triton_bb_fp8_gemm( triton.cdiv(N, META["BLOCK_SIZE_N"]), ) # noqa - fp8_gemm_bb_kernel[grid]( + fp8_blockwise_gemm_kernel[grid]( a, b, c, @@ -91,93 +91,3 @@ def triton_bb_fp8_gemm( 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, - 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) - - -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, - BLOCK_SIZE_M=64, - BLOCK_SIZE_N=64, - ) - return c diff --git a/linghe/gemm/fp32_gemm.py b/linghe/gemm/fp32_gemm.py index ed83ef8..6e0dac6 100644 --- a/linghe/gemm/fp32_gemm.py +++ b/linghe/gemm/fp32_gemm.py @@ -3,24 +3,15 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +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, @@ -70,15 +61,15 @@ def triton_fp32_gemm(x: torch.Tensor, w: torch.Tensor): assert x.is_contiguous() and w.is_contiguous() M, K = x.size() N, K = w.size() - assert M % 32 == 0 and K % 128 == 0 and N % 16 == 0 + 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 = 128 - BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = max([x for x in [16, 32, 64, 128] if N % x == 0]) + 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]( @@ -97,7 +88,6 @@ def triton_fp32_gemm(x: torch.Tensor, w: torch.Tensor): return c -# @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit def fp32_gemm_for_backward_kernel( a_ptr, @@ -150,11 +140,11 @@ def triton_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): 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, 128] if K % x == 0]) - BLOCK_SIZE_M = 32 - BLOCK_SIZE_N = 128 + 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 + num_stages = 3 fp32_gemm_for_backward_kernel[grid]( y, w, @@ -171,7 +161,6 @@ def triton_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): return c -# @triton.autotune(configs=fp32_gemm_configs, key=["M", "N", "K"]) @triton.jit def fp32_gemm_for_update_kernel( a_ptr, @@ -224,9 +213,9 @@ def triton_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): triton.cdiv(M, META["BLOCK_SIZE_M"]), triton.cdiv(N, META["BLOCK_SIZE_N"]), ) # noqa - BLOCK_SIZE_K = 128 - BLOCK_SIZE_M = max([x for x in [16, 32] if M % x == 0]) - BLOCK_SIZE_N = 128 + BLOCK_SIZE_K = max([x for x in [32, 64, 128] if K % x == 0]) + BLOCK_SIZE_M = 32 + BLOCK_SIZE_N = 64 num_warps = 4 num_stages = 3 fp32_gemm_for_update_kernel[grid]( @@ -245,6 +234,23 @@ def triton_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): 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 split_fp32_gemm_kernel( a_ptr, @@ -253,14 +259,14 @@ def split_fp32_gemm_kernel( 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, - SPLIT_COUNT: tl.constexpr, ): - pid_m = tl.program_id(axis=0) - pid_n = tl.program_id(axis=1) - pid_k = tl.program_id(axis=2) + 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) @@ -300,25 +306,28 @@ def triton_split_fp32_gemm(x: torch.Tensor, w: torch.Tensor): 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 = 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 + # 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"]), - SPLIT_COUNT, ) # noqa - num_warps = 4 - num_stages = 3 split_fp32_gemm_kernel[grid]( x, w, @@ -326,17 +335,19 @@ def triton_split_fp32_gemm(x: torch.Tensor, w: torch.Tensor): M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, SPLIT_COUNT, - num_warps=num_warps, - num_stages=num_stages, + # 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.autotune( + configs=split_fp32_gemm_configs, key=["M", "N", "K"], reset_to_zero=["c_ptr"] +) @triton.jit def split_fp32_gemm_for_backward_kernel( a_ptr, @@ -345,14 +356,14 @@ def split_fp32_gemm_for_backward_kernel( 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, - SPLIT_COUNT: tl.constexpr, ): - pid_m = tl.program_id(axis=0) - pid_n = tl.program_id(axis=1) - pid_k = tl.program_id(axis=2) + 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) @@ -392,23 +403,30 @@ def triton_split_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): assert y.is_contiguous() and w.is_contiguous() M, K = y.size() K, N = w.size() - 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 - SPLIT_COUNT = min(triton.cdiv(K, 2048), 8) + + 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"]), - SPLIT_COUNT, ) # noqa - num_warps = 4 - num_stages = 2 split_fp32_gemm_for_backward_kernel[grid]( y, w, @@ -416,35 +434,37 @@ def triton_split_fp32_gemm_for_backward(y: torch.Tensor, w: torch.Tensor): M, N, K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, SPLIT_COUNT, - num_warps=num_warps, - num_stages=num_stages, + # 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=fp32_gemm_configs, key=["M", "N", "K"]) +@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, - M, + K, + M: tl.constexpr, 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, - SPLIT_COUNT: tl.constexpr, ): - pid_m = tl.program_id(axis=0) - pid_n = tl.program_id(axis=1) - pid_k = tl.program_id(axis=2) + 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) @@ -486,34 +506,41 @@ def triton_split_fp32_gemm_for_update(y: torch.Tensor, x: torch.Tensor): assert y.is_contiguous() and x.is_contiguous() K, M = y.size() K, N = x.size() - BLOCK_SIZE_K = 64 - BLOCK_SIZE_M = max([x for x in [16, 32, 64] if M % x == 0]) - BLOCK_SIZE_N = 128 - SPLIT_COUNT = min(triton.cdiv(K, 2048), 8) + + 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"]), - SPLIT_COUNT, ) # noqa - - num_warps = 2 - num_stages = 3 split_fp32_gemm_for_update_kernel[grid]( y, x, c, + K, M, N, - K, - BLOCK_SIZE_K, - BLOCK_SIZE_M, - BLOCK_SIZE_N, SPLIT_COUNT, - num_warps=num_warps, - num_stages=num_stages, + # 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 1fe1905..594126d 100644 --- a/linghe/quant/block.py +++ b/linghe/quant/block.py @@ -24,12 +24,11 @@ def block_quant_kernel( 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, used for blockwise recipe for weight in megatron Args: @@ -43,23 +42,28 @@ def triton_block_quant(x, block_size=128, round_scale=False): """ assert x.is_contiguous() M, N = x.size() - y = torch.empty((M, N), dtype=torch.float8_e4m3fn, device=x.device) - s = torch.empty( - M // block_size, N // block_size, dtype=torch.float32, device=x.device - ) + 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, + out, + scale, M, N, BLOCK_SIZE=block_size, ROUND=round_scale, - num_stages=6, - num_warps=8, + num_stages=3, + num_warps=4, ) - return y, s + return out, scale @triton.jit diff --git a/linghe/quant/group.py b/linghe/quant/group.py index effe392..1c4d9ce 100644 --- a/linghe/quant/group.py +++ b/linghe/quant/group.py @@ -22,7 +22,7 @@ def group_quant_kernel( offs = pid * N + tl.arange(0, K * BLOCK_SIZE) n = tl.cdiv(N, K * BLOCK_SIZE) soffs = pid * (N // BLOCK_SIZE) + tl.arange(0, K) - for i in tl.range(n, flatten=True): + 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) @@ -56,7 +56,7 @@ def triton_group_quant(x, dtype=torch.float8_e4m3fn, group_size=128, round_scale 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 + grid = (M,) group_quant_kernel[grid]( x, y, s, N, group_size, K, round_scale, num_stages=5, num_warps=4 ) 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 9c96012..092655c 100644 --- a/linghe/quant/smooth.py +++ b/linghe/quant/smooth.py @@ -7,9 +7,6 @@ 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( @@ -17,7 +14,6 @@ def tokenwise_smooth_quant_kernel( q_ptr, ss_ptr, qs_ptr, - max_ptr, M, T, N: tl.constexpr, @@ -25,16 +21,12 @@ def tokenwise_smooth_quant_kernel( EVEN: tl.constexpr, REVERSE: tl.constexpr, ROUND: tl.constexpr, - CALIBRATE: 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: @@ -54,18 +46,13 @@ def tokenwise_smooth_quant_kernel( + tl.arange(0, N)[None, :], mask=indices[:, None] < M, ).to(tl.float32) - if CALIBRATE: - output_maxs = tl.maximum(tl.abs(x), output_maxs) 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 @@ -92,9 +79,6 @@ def tokenwise_smooth_quant_kernel( 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) @triton.jit @@ -103,7 +87,6 @@ def blockwise_smooth_quant_kernel( q_ptr, ss_ptr, qs_ptr, - max_ptr, M, N, H: tl.constexpr, @@ -111,10 +94,8 @@ def blockwise_smooth_quant_kernel( EVEN: tl.constexpr, REVERSE: tl.constexpr, ROUND: tl.constexpr, - CALIBRATE: 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) @@ -127,9 +108,6 @@ def blockwise_smooth_quant_kernel( 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) if REVERSE: x = x * smooth_scale else: @@ -172,15 +150,9 @@ def blockwise_smooth_quant_kernel( def triton_smooth_quant( - x, - smooth_scale, - x_q=None, - x_scale=None, - reverse=False, - round_scale=False, - calibrate=False, + x, smooth_scale, x_q=None, x_scale=None, reverse=False, round_scale=False ): - """""" + # it may be used for sharded weight quantization, therefore M is not exact batch size M, N = x.shape device = x.device if x_q is None: @@ -190,20 +162,13 @@ def triton_smooth_quant( if triton.next_power_of_2(N) == N and N <= 8192: W = 8192 // N T = 8 - # it may used in shard weight quantization, therefore M is not batch size - # assert M % (W * T) == 0 EVEN = M % (W * T) == 0 g = triton.cdiv(M, W * T) - if calibrate: - x_maxs = torch.empty((g, N), device=device, dtype=torch.bfloat16) - else: - x_maxs = None tokenwise_smooth_quant_kernel[(g,)]( x, x_q, smooth_scale, x_scale, - x_maxs, M, T, N, @@ -211,28 +176,21 @@ def triton_smooth_quant( EVEN, reverse, round_scale, - calibrate, num_stages=3, num_warps=4, ) - if calibrate: - x_maxs = x_maxs.amax(0).float() else: - H = max([x for x in [256, 512, 1024, 2048] if N % x == 0]) + # N may be 576 in MLA + H = max([x for x in [64, 128, 256, 512, 1024, 2048] if N % x == 0]) W = 8 if M > 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, @@ -240,359 +198,237 @@ def triton_smooth_quant( EVEN, reverse, round_scale, - calibrate, num_stages=3, 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( +def transpose_smooth_quant_kernel( x_ptr, q_ptr, ss_ptr, qs_ptr, - subrow_scales_ptr, - tail_ri, - tail_si, - head_ri, - head_ei, - size, + M, N, + P, + H: tl.constexpr, W: tl.constexpr, - TAIL: tl.constexpr, - HEAD: tl.constexpr, + EVEN: 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, + 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 - 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) + scale = tl.maximum(x_max / 448.0, 1e-30) + if ROUND: + scale = tl.exp2(tl.ceil(tl.log2(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( + 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 ) - x *= smooth_scale - x /= scale - xq = tl.minimum(tl.maximum(x, -448), 448) + 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 + head_ri * N + i * W + tl.arange(0, W), - xq.to(q_ptr.dtype.element_ty), - mask=mask, + q_ptr + toffs, tl.trans(x), mask=(i * H + tl.arange(0, H)[None, :] < P) ) + offs += H * N + toffs += H + soffs += H -def triton_subrow_smooth_quant( - x, - smooth_scale, - x_q, - x_scale, - subrow_scales, - offset, - size, - reverse=False, - round_scale=False, +def triton_transpose_smooth_quant( + x, smooth_scale, reverse=False, pad=True, 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 - - if (offset + size) % N == 0: - head_ri = 0 - head_ei = 0 # head_size = head_ei - HEAD = False - else: - head_ri = (offset + size) // N - head_ei = (offset + size) % N - HEAD = True + # M should be padded to mutiple of 32 if pad is True + 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 = (1,) - subrow_smooth_quant_kernel[grid]( + grid = (triton.cdiv(N, W),) + transpose_smooth_quant_kernel[grid]( x, x_q, smooth_scale, x_scale, - subrow_scales, - tail_ri, - tail_si, - head_ri, - head_ei, - size, + M, N, + P, + H, W, - TAIL, - HEAD, + EVEN, reverse, round_scale, num_stages=3, - num_warps=1, + num_warps=4 if N >= 8192 else 4, ) + return x_q, x_scale @triton.jit -def depracated_tokenwise_smooth_quant_kernel( +def batch_smooth_quant_kernel( x_ptr, q_ptr, ss_ptr, qs_ptr, - M, - W, + count_ptr, + accum_ptr, + T, 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)) + eid = tl.program_id(axis=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 - 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) + count = tl.load(count_ptr + eid) + ei = tl.load(accum_ptr + eid) + si = ei - count - scale = x_max / 448.0 + 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) + x *= smooth_scale + scale = tl.maximum(tl.max(tl.abs(x)) / 448.0, 1e-30) if ROUND: scale = tl.exp2(tl.ceil(tl.log2(scale))) - tl.store(qs_ptr + pid * W + i, scale, mask=pid * W + i < M) - x /= 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 + pid * W * N + i * N + tl.arange(0, N), xq, mask=pid * W + i < M - ) + tl.store(q_ptr + si * N + i * N + tl.arange(0, N), xq) -def triton_depracated_tokenwise_smooth_quant( - x, smooth_scale, x_q=None, x_scale=None, reverse=False, round_scale=False +def triton_batch_smooth_quant( + x, smooth_scales, token_count_per_expert, reverse=False, round_scale=False ): - """""" - # row-wise read, row-wise write + """ + 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 - 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]( + n_expert = token_count_per_expert.shape[0] + 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 + + grid = (n_expert, T) + batch_smooth_quant_kernel[grid]( x, x_q, - smooth_scale, + smooth_scales, x_scale, - M, - W, + token_count_per_expert, + accum_token_count, + T, N, reverse, round_scale, num_stages=3, - num_warps=8, + num_warps=4, ) return x_q, x_scale @triton.jit -def batch_smooth_quant_kernel( +def batch_transpose_smooth_quant_kernel( x_ptr, q_ptr, ss_ptr, qs_ptr, - xm_ptr, count_ptr, - accum_ptr, - T, - N: tl.constexpr, + N, + H: tl.constexpr, + W: tl.constexpr, + E: 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 + eid = tl.program_id(axis=0) + bid = tl.program_id(axis=1) - # row-wise read, row-wise write - smooth_scale = tl.load(ss_ptr + i_expert * N + tl.arange(0, N)) - if not REVERSE: - smooth_scale = 1.0 / smooth_scale + count = tl.load(count_ptr + eid) + round_count = tl.cdiv(count, 32) * 32 - if CALIBRATE: - x_maxs = tl.zeros((N,), dtype=tl.float32) + counts = tl.load(count_ptr + tl.arange(0, E)) + si = tl.sum(tl.where(tl.arange(0, E) < eid, counts, 0)) - count = tl.load(count_ptr + i_expert) - ei = tl.load(accum_ptr + i_expert) - 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)): - 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: - 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) - - 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, -): - """""" - 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) - 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) - - grid = (128,) - batch_smooth_quant_kernel[grid]( - x, - x_q, - smooth_scales, - x_scale, - x_maxs, - token_count_per_expert, - accum_token_count, - T, - N, - reverse, - round_scale, - calibrate, - num_stages=3, - num_warps=8, - ) - if calibrate: - x_maxs = x_maxs.view(n_expert, T, N).amax(1) - return x_q, x_scale, x_maxs - - -@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, -): - 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)) + 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 @@ -601,14 +437,13 @@ def batch_pad_transpose_smooth_quant_kernel( + si * N + i * H * N + bid * W - + tl.arange(0, H)[:, None] + + 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))) @@ -616,9 +451,8 @@ def batch_pad_transpose_smooth_quant_kernel( 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 @@ -627,7 +461,7 @@ def batch_pad_transpose_smooth_quant_kernel( + si * N + i * H * N + bid * W - + tl.arange(0, H)[:, None] + + tl.arange(0, H)[:, None] * N + tl.arange(0, W)[None, :], mask=indices[:, None] < count, ).to(tl.float32) @@ -636,10 +470,10 @@ def batch_pad_transpose_smooth_quant_kernel( xq = tl.trans(x.to(q_ptr.dtype.element_ty)) tl.store( q_ptr - + bias * N + + round_si * N + bid * W * round_count + i * H - + tl.arange(0, W)[:, None] + + tl.arange(0, W)[:, None] * round_count + tl.arange(0, H)[None, :], xq, mask=indices[None, :] < round_count, @@ -658,38 +492,32 @@ def batch_pad_transpose_smooth_quant_kernel( """ -def triton_batch_pad_transpose_smooth_quant( +def triton_batch_transpose_smooth_quant( x, smooth_scales, token_count_per_expert, splits, - x_q=None, - x_scale=None, - x_maxs=None, + 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, @@ -702,127 +530,6 @@ def triton_batch_pad_transpose_smooth_quant( 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, - ) - return x_q, x_scale - - @triton.jit def transpose_rescale_smooth_quant_kernel( x_ptr, @@ -840,7 +547,6 @@ def transpose_rescale_smooth_quant_kernel( 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) @@ -937,7 +643,7 @@ def triton_transpose_rescale_smooth_quant( 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 @@ -967,150 +673,134 @@ def triton_transpose_rescale_smooth_quant( return xt_q, x_scale -""" -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 - -""" +@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 -""" -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 + 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, + ) -transpose: transpose quantized x for wgrad -pad: # pad M to be multiplier of 32, including quant scales and transposed x + 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, + ) -# y = x @ w -# dx = y @ wT -# dwT = yT @ x -def triton_smooth_quant_input( +def triton_subrow_smooth_quant( 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 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, + 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 - - -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 + tail_ri = offset // N + tail_si = offset % N + TAIL = True - 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, - ) + head_ri = (offset + size) // N + head_ei = (offset + size) % N + HEAD = True - # 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, - ) + 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/benchmark.py b/linghe/tools/benchmark.py index 71ea46d..e44fe10 100644 --- a/linghe/tools/benchmark.py +++ b/linghe/tools/benchmark.py @@ -45,7 +45,8 @@ def benchmark_func( ProfilerActivity.CPU, ProfilerActivity.CUDA, ProfilerActivity.XPU, - ] + ], + with_stack=True, ) as prof: for i in range(n_profile): fn(*args, **kwargs) diff --git a/linghe/tools/check.py b/linghe/tools/check.py index 460c2e3..7755590 100644 --- a/linghe/tools/check.py +++ b/linghe/tools/check.py @@ -91,6 +91,7 @@ def output_check( # 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() @@ -98,6 +99,9 @@ def output_check( 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]) @@ -109,7 +113,8 @@ def output_check( 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 opt: {opt_str} \n" + 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 diff --git a/linghe/tools/util.py b/linghe/tools/util.py index 33edbff..cb368d0 100644 --- a/linghe/tools/util.py +++ b/linghe/tools/util.py @@ -6,6 +6,8 @@ import math import torch +import triton +import triton.language as tl def round_up(x, b=16): @@ -66,6 +68,12 @@ 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 @@ -82,6 +90,13 @@ def torch_blockwise_quant(x, round_scale=True, padding=False): 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() @@ -98,9 +113,15 @@ def torch_block_quant(w, B=128, dtype=torch.float8_e4m3fn, round_scale=False): return wq, scale -def torch_mxfp8_quant(x): +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 + m, N = x.shape # current m is multiple of 32 assert N % 128 == 0 if m % 128 != 0: M = (m + 127) // 128 * 128 @@ -111,7 +132,9 @@ def torch_mxfp8_quant(x): 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] + 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) @@ -119,14 +142,43 @@ def torch_mxfp8_quant(x): 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: @@ -138,7 +190,7 @@ def torch_smooth_quant(x, smooth_scale, reverse=False, round_scale=False): 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( @@ -514,3 +566,65 @@ def torch_fp32_scaler_scaled_mm(x, weight, x_scale, weight_scale): 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/tests/test_add.py b/tests/test_add.py index e6babd7..40264fd 100644 --- a/tests/test_add.py +++ b/tests/test_add.py @@ -3,57 +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.check import output_check -from linghe.utils.add import triton_inplace_add +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 +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" - outputs = torch.randn(M, N, dtype=dtype, device=device) - x = torch.randn(M, N, dtype=dtype, device=device) - - out = outputs.clone() - triton_inplace_add(out, x) - out_ref = outputs + x - output_check(out_ref, out, "sum") - - if bench: - n_repeat = 100 - 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, - ) - - 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, - ) - - -if __name__ == "__main__": - test_triton_inplace_add(M=4096, N=4096) + 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") + + 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 + ) + + +@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" + + 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)] + + 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") + + 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 611cf72..243d531 100644 --- a/tests/test_blockwise_fp8_gemm.py +++ b/tests/test_blockwise_fp8_gemm.py @@ -3,14 +3,20 @@ 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.tools.benchmark import benchmark_func +from linghe.gemm.blockwise_fp8_gemm import triton_blockwise_fp8_gemm from linghe.tools.check 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" B = 64 @@ -32,64 +38,22 @@ def test_triton_bb_gemm(M=4096, N=4096, K=4096, bench=False): ) 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) + 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) - 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.to(dtype), y, "y", atol=1.0, rtol=0.05) - - 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=False) - test_triton_tt_gemm(M=4096, N=8192, K=2048, bench=False) + 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 index 0e726ef..f97a749 100644 --- a/tests/test_blockwise_quant.py +++ b/tests/test_blockwise_quant.py @@ -1,3 +1,4 @@ +import pytest import torch from linghe.quant.block import ( @@ -5,7 +6,6 @@ triton_blockwise_quant, triton_batch_blockwise_quant, ) -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.tools.util import ( torch_block_quant, @@ -43,7 +43,13 @@ def torch_batch_blockwise_quant(x, token_count_per_expert_list, round_scale=True return q_ref, s_ref, qt_ref, st_ref -def test_block_quant(M=8192, N=4096, bench=False): +@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 @@ -52,11 +58,16 @@ def test_block_quant(M=8192, N=4096, bench=False): output_check(x_q_ref.float(), x_q.float(), "data") output_check(x_s_ref.float(), x_s.float(), "scale") - if bench: - benchmark_func(triton_block_quant, x, round_scale=True, ref_bytes=M * N * 4) + benchmark(triton_block_quant, x, round_scale=True, ref_bytes=M * N * 4) -def test_blockwise_quant(M=8192, N=4096, bench=False): +@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 @@ -69,11 +80,16 @@ def test_blockwise_quant(M=8192, N=4096, bench=False): output_check(xt_q_ref.float(), xt_q.float(), "t.data") output_check(xt_s_ref.float(), xt_s.float(), "t.scale") - if bench: - benchmark_func(triton_blockwise_quant, x, round_scale=True, ref_bytes=M * N * 4) + benchmark(triton_blockwise_quant, x, round_scale=True, ref_bytes=M * N * 4) -def test_batch_block_quant(M=16384, N=2048, n_experts=32, topk=2, bench=False): +@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 @@ -98,18 +114,11 @@ def test_batch_block_quant(M=16384, N=2048, n_experts=32, topk=2, bench=False): 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") - if bench: - benchmark_func( - triton_batch_blockwise_quant, - x, - token_count_per_expert, - token_count_per_expert_list, - round_scale=True, - ref_bytes=M * N * 4, - ) - - -if __name__ == "__main__": - test_block_quant(M=8192, N=4096, bench=False) - test_blockwise_quant(M=8192, N=4096, bench=False) - test_batch_block_quant(M=16384, N=2048, n_experts=32, topk=2, bench=False) + 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 96943c3..47d15d2 100644 --- a/tests/test_channel_quant.py +++ b/tests/test_channel_quant.py @@ -3,6 +3,7 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch from linghe.quant.channel import ( @@ -10,12 +11,22 @@ triton_row_quant, triton_tokenwise_row_quant, ) -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.tools.util import torch_row_quant -def test_row_quant(M=4096, N=4096, round_scale=True, bench=False): +@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 @@ -30,30 +41,19 @@ def test_row_quant(M=4096, N=4096, round_scale=True, bench=False): output_check(x_q_ref, x_q, name="data") output_check(x_scale_ref, x_scale, name="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) + 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 7bb2b0d..8d5417d 100644 --- a/tests/test_channelwise_fp8_gemm.py +++ b/tests/test_channelwise_fp8_gemm.py @@ -3,10 +3,10 @@ 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.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.utils.add import triton_inplace_add @@ -28,7 +28,13 @@ 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" @@ -45,94 +51,89 @@ def test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False): output_check(y_ref, y, name="y", atol=-1) - if bench: - y_bf16 = torch.randn(M, N, dtype=dtype, device=device) - y_fp16 = y.to(torch.float16) - y_fp32 = y.to(torch.float32) + 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 + 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, - ) + 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, + ) -if __name__ == "__main__": - test_triton_channelwise_gemm(M=4096, N=4096, K=4096, bench=False) + 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 index 52afe0f..516b142 100644 --- a/tests/test_dist_loss.py +++ b/tests/test_dist_loss.py @@ -17,10 +17,8 @@ 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( @@ -186,7 +184,7 @@ def test_triton_softmax_cross_entropy( ) 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=False + M=8192, N=157184, coef=1.0, grad_coef=1.0, inplace=False, group=pg, bench=True ) test_triton_softmax_cross_entropy( M=8192, diff --git a/tests/test_embedding.py b/tests/test_embedding.py index be9818c..89ec033 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -3,10 +3,10 @@ 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.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.utils.emb import ( triton_embedding_forward, @@ -17,35 +17,81 @@ ) -def test_scan(M=4096, bench=False): +@pytest.mark.parametrize( + "B,M", + [ + (None, 8192), + (None, 4097), + (4, 8192), + (2, 4097), + ], +) +def test_scan(B, M, benchmark): device = "cuda: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") + if B is None or B == 0: + input_ids = torch.randint(0, 10000, (M,), dtype=torch.int32, device=device) - if bench: - ref_time = benchmark_func(triton_scan_and_count, sorted_ids) + 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") -def test_embedding(B=2, M=4096, V=150000, D=4096, transpose=False, bench=False): + 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 - weights.grad = torch.zeros((V, D), dtype=dtype, device=device) + 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: @@ -53,125 +99,102 @@ def test_embedding(B=2, M=4096, V=150000, D=4096, transpose=False, bench=False): else: dy = torch.randn((B, M, D), device=device, dtype=dtype) y_ref.backward(dy, retain_graph=True) - grad_ref = weights.grad.clone().detach() - grad = weights.grad - grad.zero_() - y = triton_embedding_forward(input_ids, weights.data_ptr(), D, dtype) - output_check(y_ref, y, name="y") - - triton_embedding_backward(dy, input_ids, grad.data_ptr(), grad.dtype) - output_check(grad_ref, grad.to(dtype), name="grad") - - grad.zero_() + 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") - if bench: - ref_bytes = B * M * D * 4 - ref_time = benchmark_func(embedding.forward, input_ids, ref_bytes=ref_bytes) - benchmark_func( - embedding_lookup, input_ids, weights, ref_time=ref_time, ref_bytes=ref_bytes - ) + 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_func( - y_ref.backward, dy, retain_graph=True, ref_bytes=ref_bytes - ) - benchmark_func( - triton_atomic_embedding_backward, - dy, - input_ids, - grad.data_ptr(), - grad.dtype, - ref_time=ref_time, - ref_bytes=ref_bytes, - ) - benchmark_func( - triton_sync_embedding_backward, - dy, - input_ids, - grad.data_ptr(), - grad.dtype, - ref_time=ref_time, - ref_bytes=ref_bytes, - ) - benchmark_func( - triton_embedding_backward, - dy, - input_ids, - grad.data_ptr(), - grad.dtype, - ref_time=ref_time, - ref_bytes=ref_bytes, - ) - benchmark_func( - y.backward, dy, retain_graph=True, 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) -def test_fused_embedding( - B=2, M=4096, V=150000, D=4096, use_main_grad=True, transpose=False, bench=False -): +@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" - grad_name = "main_grad" if use_main_grad else "grad" 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 - weights.grad = torch.zeros((V, D), dtype=dtype, device=device) - if use_main_grad: - weights.main_grad = torch.zeros((V, D), dtype=dtype, device=device) - grad = weights.main_grad - else: - grad = weights.grad + 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) - grad_ref = weights.grad.clone().detach() - - grad.zero_() - y = triton_embedding_forward(input_ids, weights.data_ptr(), D, dtype) - output_check(y_ref, y, name="y") - - triton_embedding_backward(dy, input_ids, grad.data_ptr(), grad.dtype) - output_check(grad_ref, grad.to(dtype), name="grad") - grad.zero_() + 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, grad.to(dtype), name="grad") - - if bench: - ref_bytes = B * M * D * 4 - ref_time = benchmark_func(embedding.forward, input_ids) - benchmark_func( - fused_accumulation_embedding_lookup, - input_ids, - weights, - grad_name=grad_name, - ref_time=ref_time, - ref_bytes=ref_bytes, - ) - - ref_time = benchmark_func(y_ref.backward, dy, retain_graph=True) - benchmark_func( - y.backward, dy, retain_graph=True, ref_time=ref_time, ref_bytes=ref_bytes - ) + 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, + ) -if __name__ == "__main__": - test_scan(M=8192, bench=False) - test_embedding(B=1, M=8192, V=150000, D=8192, transpose=False, bench=False) - test_embedding(B=2, M=4096, V=150000, D=4096, transpose=True, bench=False) - test_fused_embedding(B=1, M=8192, V=150000, D=8192, transpose=False, bench=False) - test_fused_embedding(B=1, M=4096, V=150000, D=8192, transpose=False, bench=False) - test_fused_embedding(B=2, M=4096, V=150000, D=8192, transpose=True, bench=False) - test_fused_embedding(B=0, M=4096, V=150000, D=8192, transpose=True, bench=False) + 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 6fea714..1333971 100644 --- a/tests/test_fp32_gemm.py +++ b/tests/test_fp32_gemm.py @@ -3,8 +3,12 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest + import torch +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, @@ -14,7 +18,7 @@ triton_split_fp32_gemm_for_backward, triton_split_fp32_gemm_for_update, ) -from linghe.tools.benchmark import benchmark_func +from linghe.experimental.gemm import triton_tma_persistent_matmul from linghe.tools.check import output_check @@ -36,12 +40,23 @@ 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): +@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" x = torch.randn(M, K, dtype=dtype, device=device, requires_grad=True) - w = torch.randn(N, 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) @@ -57,98 +72,191 @@ def test_fp32_matmul(M=2048, N=256, K=8192, bench=False): 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) - y = triton_split_fp32_gemm(x, w) + 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-2) + output_check(dw_ref, dw.to(dtype), name="split.dw", atol=2e-1, rtol=2e-1) - x.grad = None - w.grad = None - y = fp32_gemm(x, w) - y.backward(gradient=dy) - dx = x.grad - dw = w.grad - 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) + 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) - if bench: - ref_bytes = M * K * 6 + N * K * 6 + M * N * 4 - ref_flops = 2 * M * N * K - ref_time = benchmark_func( - torch_fp32_matmul, x, w, ref_bytes=ref_bytes, ref_flops=ref_flops - ) - benchmark_func( - triton_fp32_gemm, - x, - w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ref_time=ref_time, - ) - benchmark_func( - triton_split_fp32_gemm, - x, - w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ref_time=ref_time, - ) - - ref_bytes = M * K * 10 + N * K * 4 + M * N * 4 - ref_time = benchmark_func( - torch_fp32_matmul_backward, - dy, - w.float(), - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ) - benchmark_func( - triton_fp32_gemm_for_backward, - dy, - w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ref_time=ref_time, - ) - benchmark_func( - triton_split_fp32_gemm_for_backward, - dy, - w, - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ref_time=ref_time, - ) - - ref_bytes = M * K * 4 + N * K * 12 + M * N * 4 - ref_time = benchmark_func( - torch_fp32_matmul_update, - dy, - x.float(), - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ) - benchmark_func( - triton_fp32_gemm_for_update, - dy, - x, - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ref_time=ref_time, - ) - benchmark_func( - triton_split_fp32_gemm_for_update, - dy, - x, - ref_bytes=ref_bytes, - ref_flops=ref_flops, - ref_time=ref_time, - ) - - -def test_BMK_fp32_matmul(B=2, M=2048, N=256, K=8192, bench=False): + 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" @@ -165,39 +273,29 @@ def test_BMK_fp32_matmul(B=2, M=2048, N=256, K=8192, bench=False): x.grad = None w.grad = None - y = fp32_gemm(x, w) + 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=1e-1, rtol=2e-2) - output_check(dw_ref, dw, name="update", atol=1e-1, rtol=2e-2) - - 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_flops=2 * M * N * K, - ) - benchmark_func( - fp32_gemm, - x, - w, - 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, - ) - - -if __name__ == "__main__": - test_fp32_matmul(M=4096, N=256, K=8192, bench=False) - test_fp32_matmul(M=16384, N=256, K=2048, bench=False) - test_fp32_matmul(M=128, N=16, K=128, bench=False) - test_BMK_fp32_matmul(B=2, M=2048, N=16, K=8192, bench=False) - test_BMK_fp32_matmul(B=2, M=2048, N=256, K=8192, bench=False) - test_BMK_fp32_matmul(B=2, M=128, N=16, K=128, bench=False) + 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 index 3cd8baf..9761a3a 100644 --- a/tests/test_gate.py +++ b/tests/test_gate.py @@ -3,29 +3,31 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch import torch.nn.functional as F -from linghe.tools.benchmark import benchmark_func 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, transpose=True + 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() - if transpose: - length, bs, dim = gate.shape - else: - bs, length, dim = gate.shape + length, bs, dim = gate.shape d = dim // group_size attn_output = x.view(bs, length, group_size, d) outputs = [] @@ -38,15 +40,56 @@ def torch_group_rms_norm_gate_forward( o = F.rms_norm(attn_output[:, :, i], [d], weight=weight, eps=eps) outputs.append(o) outputs = torch.stack(outputs, 2).view(bs, length, dim) - if transpose: - outputs = outputs.transpose(0, 1) + outputs = outputs.transpose(0, 1) gate = F.sigmoid(gate) - outputs = (outputs * gate).to(dtype) + 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, transpose=True + grad_output, x, gate, weight, eps=1e-6, group_size=4, native=True ): dtype = grad_output.dtype grad_output = grad_output.float() @@ -54,135 +97,187 @@ def torch_group_rms_norm_gate_backward( 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 + 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=1, - length=4096, - dim=4096, - group_size=4, - transpose=True, - share=False, - coef=1.0, - grad_coef=1.0, - bench=False, + bs, length, dim, group_size, contiguous, share, coef, grad_coef, native, benchmark ): dtype = torch.bfloat16 device = "cuda:0" - x = torch.randn(bs, length, dim, dtype=dtype, requires_grad=True, device=device) + 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 transpose: + if contiguous: gate = ( torch.randn(length, bs, dim, dtype=dtype, device=device) * coef ).requires_grad_() - grad_output = ( - torch.randn(length, bs, dim, dtype=dtype, device=device) * grad_coef - ) else: - gate = ( - torch.randn(bs, length, dim, dtype=dtype, device=device) * coef - ).requires_grad_() - grad_output = ( - torch.randn(bs, length, dim, dtype=dtype, device=device) * grad_coef - ) + 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, transpose=transpose - ) - output = triton_group_rms_norm_gate_forward( - x, gate, weight, group_size=group_size, transpose=transpose + 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_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 + 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") - if bench: - benchmark_func( - torch_group_rms_norm_gate_forward, - x, - gate, - weight, - group_size=group_size, - transpose=transpose, - ref_bytes=bs * length * dim * 6, + 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 + ) - benchmark_func( - triton_group_rms_norm_gate_forward, - x, - gate, - weight, - group_size=group_size, - transpose=transpose, - ref_bytes=bs * length * dim * 6, + 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 ) + ) - benchmark_func( - triton_group_rms_norm_gate_backward, - grad_output, - x, - gate, - weight, - group_size=group_size, - transpose=transpose, - ref_bytes=bs * length * dim * 10, + 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, + ) -if __name__ == "__main__": - 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=2, - length=4096, - dim=2048, - group_size=4, - transpose=False, - share=True, - bench=False, - ) - test_group_rms_norm_gate(bs=1, length=4096, dim=4096, group_size=4, bench=False) - test_group_rms_norm_gate( - bs=2, length=4096, dim=1536, group_size=4, transpose=True, bench=False - ) - test_group_rms_norm_gate( - bs=2, - length=4096, - dim=1536, - group_size=4, - transpose=False, - coef=10000.0, - grad_coef=10000.0, - bench=False, - ) - test_group_rms_norm_gate( - bs=2, - length=4096, - dim=1536, - group_size=4, - transpose=False, - coef=0.0, - grad_coef=0.0, - bench=False, + 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 c47aeef..032aa62 100644 --- a/tests/test_gather.py +++ b/tests/test_gather.py @@ -3,27 +3,32 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch from linghe.quant.block import triton_batch_blockwise_quant -from linghe.tools.benchmark import benchmark_func +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_index_select, + triton_permute_with_indices, triton_permute_with_mask_map, - triton_smooth_permute_with_indices, - triton_smooth_permute_with_mask_map, - triton_smooth_weighted_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_batch_block_pad_permute_with_indices, + triton_batch_mxfp8_permute_with_indices, + triton_make_chunk_sort_map, ) @@ -65,48 +70,104 @@ def torch_scatter(logits, routing_map, weights): logits[routing_map] = weights -# optional dequant and smooth and quant +# 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 -# desmooth,dequant, gather, pad, transpose, smooth, quant +# 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_fused_permute_with_indices( x_q, x_scale, org_smooth_scale, @@ -136,7 +197,7 @@ def torch_batch_transpose_smooth_permute_with_indices( 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_q, y_scale = torch_smooth_quant( y.t().contiguous(), smooth_scale, reverse=True, round_scale=round_scale ) scale_refs.append(y_scale.view(-1)) @@ -197,7 +258,61 @@ def torch_batch_block_pad_permute_with_indices( return q_ref, s_ref, qt_ref, st_ref, probs_refs -def test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): +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" @@ -215,66 +330,35 @@ def test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False): _, 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 + 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", ) - probs, mask_map, token_count_per_expert, indices, row_id_map = torch_make_indices( - logits, topk=topk, bias=0.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 ) - - 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, + row_id_map_torch, resort_row_id_map_torch = torch_make_chunk_sort_map( + num_global_tokens_per_local_expert ) - 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, - ) + assert torch.equal(row_id_map_torch, row_id_map) + assert torch.equal(resort_row_id_map_torch, resort_row_id_map) -def test_triton_permute_with_mask_map( - M=4096, N=4096, n_experts=256, topk=8, bench=False -): +@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 @@ -288,7 +372,7 @@ def test_triton_permute_with_mask_map( ) 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") @@ -331,55 +415,61 @@ def test_triton_permute_with_mask_map( 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, - tokens_per_expert=token_count_per_expert, - n_repeat=n_repeat, - ref_time=ref_time, - ref_bytes=ref_bytes, - ) + 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, + ) -def test_triton_smooth_permute_with_mask_map( - M=4096, N=4096, n_experts=32, topk=8, round_scale=True, bench=False +@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 @@ -394,101 +484,174 @@ def test_triton_smooth_permute_with_mask_map( 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=dtype, 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, + 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 = triton_smooth_permute_with_indices( - grad_data, - grad_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") - # 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, + benchmark( + triton_batch_smooth_permute_with_indices, + grad_output, smooth_scales, - reverse=False, + 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(q_ref.float(), permuted_data.float(), name="smoothed.data", rtol=0.125) - output_check(scale_ref.float(), permuted_scale.float(), "smoothed.scale") + 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, + ) + - q_ref, scale_ref = torch_smooth_permute_with_indices( +@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 + 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 + ) + + token_count_per_expert_list = token_count_per_expert.tolist() + out_tokens = sum(token_count_per_expert_list) + + 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, - None, + grad_scale, indices, + org_smooth_scales, smooth_scales, token_count_per_expert_list, round_scale=round_scale, ) - permuted_data, permuted_scale = triton_smooth_permute_with_mask_map( + y_q, y_scale = triton_batch_smooth_fused_permute_with_indices( grad_data, - row_id_map, - None, - M, - n_experts, - out_tokens, - N, + 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.float(), permuted_data.float(), name="smoothed.data", rtol=0.125) - 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, - ) + 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, + ) -def test_triton_batch_transpose_smooth_permute_with_indices( - M=1024, N=2048, n_experts=32, topk=8, bench=False +@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: @@ -511,7 +674,6 @@ def test_triton_batch_transpose_smooth_permute_with_indices( 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"] @@ -522,7 +684,7 @@ def test_triton_batch_transpose_smooth_permute_with_indices( 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_permute_with_indices( + x_q_ref, x_scale_ref = torch_batch_transpose_smooth_fused_permute_with_indices( x, scale, org_smooth_scale, @@ -532,7 +694,7 @@ def test_triton_batch_transpose_smooth_permute_with_indices( round_scale=True, ) - x_q, x_scale = triton_batch_transpose_smooth_permute_with_indices( + x_q, x_scale = triton_batch_transpose_smooth_fused_permute_with_indices( x, scale, org_smooth_scale, @@ -542,61 +704,43 @@ def test_triton_batch_transpose_smooth_permute_with_indices( token_count_per_expert_list, round_scale=True, ) - output_check(x_q_ref.float(), x_q.float(), name="smoothed.data", rtol=0.125) + output_check(x_q_ref, x_q, name="smoothed.data", rtol=0.125) output_check(x_scale_ref.float(), x_scale.float(), "smoothed.scale") - x_q_ref, x_scale_ref = torch_batch_transpose_smooth_permute_with_indices( + benchmark( + torch_batch_transpose_smooth_fused_permute_with_indices, x, - None, - None, + scale, + org_smooth_scale, smooth_scales, indices, token_count_per_expert_list, round_scale=True, + ref_bytes=out_tokens * N * 2, ) - - x_q, x_scale = triton_batch_transpose_smooth_permute_with_indices( + benchmark( + triton_batch_transpose_smooth_fused_permute_with_indices, x, - None, - None, + 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, ) - 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, - ) -def test_batch_block_pad_permute_with_indices( - M=16384, N=2048, n_experts=32, topk=2, bench=False -): +@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 @@ -633,65 +777,112 @@ def test_batch_block_pad_permute_with_indices( output_check(xt_s_ref.float(), xt_s.view(-1).float(), "t.scale") output_check(p_ref.float(), p.view(-1).float(), "prob") - if bench: - benchmark_func( - 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_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_func( - 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_func( - triton_batch_blockwise_quant, - xs, - token_count_per_expert, - token_count_per_expert_list, - 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, + ) -if __name__ == "__main__": - test_make_id_map(M=4098, n_experts=32, topk=2, bias=0.0, bench=False) +@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() - test_triton_permute_with_mask_map( - M=16384, N=2048, n_experts=32, topk=8, 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 ) - 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) + x = torch.randn((M, N), dtype=torch.bfloat16, device=device) - test_triton_batch_transpose_smooth_permute_with_indices( - M=16384, N=2048, n_experts=32, topk=2, bench=False + 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 + ) ) - test_triton_batch_transpose_smooth_permute_with_indices( - M=8192, N=4096, n_experts=32, topk=2, bench=False + + 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") - test_batch_block_pad_permute_with_indices( - M=8192 * 2, N=2048, n_experts=32, topk=2, bench=False + 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, ) - test_batch_block_pad_permute_with_indices( - M=0, N=2048, n_experts=32, topk=2, bench=False + + 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, ) - test_batch_block_pad_permute_with_indices( - M=8192, N=1536, n_experts=32, topk=2, bench=False + 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 d9d96dd..27672e2 100644 --- a/tests/test_group_quant.py +++ b/tests/test_group_quant.py @@ -3,30 +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.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): +@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, 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) - test_group_quant(M=2049, N=1536, 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_la.py b/tests/test_la.py index 32f681f..be2fa46 100644 --- a/tests/test_la.py +++ b/tests/test_la.py @@ -5,6 +5,7 @@ import math +import pytest import torch from linghe.attn.la import ( @@ -12,7 +13,6 @@ triton_lightning_attention_backward, triton_fused_lightning_attention_backward, ) -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check @@ -86,9 +86,13 @@ def make_varlen_input( return q, k, v -def test_la( - bs=1, length=4096, qo_heads=16, kv_heads=16, dim=128, digest=False, bench=False -): +@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 @@ -138,85 +142,26 @@ def test_la( 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) - if bench: - ref_bytes = bs * length * qo_heads * dim * 8 + bs * qo_heads * dim * dim * 8 - benchmark_func( - triton_lightning_attention_forward, - q, - k, - v, - decay_scales, - ref_bytes=ref_bytes, - ) - benchmark_func( - triton_lightning_attention_backward, - g, - q, - k, - v, - decay_scales, - ref_bytes=ref_bytes, - ) - benchmark_func( - triton_fused_lightning_attention_backward, - g, - q, - k, - v, - state, - decay_scales, - ref_bytes=ref_bytes, - ) - - -# def test_varlen_la(qls=[1024,1024], qo_heads=16, kv_heads=16, dim=128, digest=False, bench=False): -# device = torch.device('cuda:0') -# dtype = torch.bfloat16 -# kls = qls - -# assert all([x<=kls[i] for i,x in enumerate(qls)]) -# bs = len(qls) - -# ref_bytes = sum(qls) * qo_heads * dim * 8 + bs * qo_heads * dim * dim * 8 - -# q, k, v = make_input(qo_heads=qo_heads, kv_heads=kv_heads, dim=dim, qls=qls, kls=kls) - -# 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 = 2**(-0.5 * torch.ones(qo_head, dtype=torch.float32, device=device)) -# lengths = torch.tensor([0] + qls, device=device, dtype=torch.long) -# cu_seqlens = torch.cumsum(lengths, 0) -# padded_cu_seqlens = cu_seqlens - -# output_ref, state_ref = torch_varlen_linear_attn(q, k, v, s, decay_scales, cu_seqlens, padded_cu_seqlens) - -# max_q_length = max(qls) -# output = triton_lightning_attention_forward(q, k, v, decay_scales, cu_seqlens, padded_cu_seqlens, max_q_length) - -# output_check(output_ref, output, name='output', rtol=0.1, atol=0.1) -# output_check(state_ref, s, name='state', rtol=0.01, atol=0.01) - -# if digest: -# print( -# f"output_ref max:{torch.max(output_ref).item():.3f} min:{torch.min(output_ref).item():.3f}") -# print( -# f"output max:{torch.max(output).item():.3f} min:{torch.min(output).item():.3f}") - -# print("output_ref[:,0,0]", output_ref[:, 0, 0]) -# print("output[:,0,0]", output[:, 0, 0]) - -# print("output_ref[0,:,0]", output_ref[0, :, 0]) -# print("output[0,:,0]", output[0, :, 0]) - -# print("output_ref[0,0,:]", output_ref[0, 0, :]) -# print("output[0,0,:]", output[0, 0, :]) - -# if bench: -# benchmark_func(triton_lightning_attention_forward, q, k, v, decay_scales, cu_seqlens, padded_cu_seqlens, max_q_length, ref_bytes=ref_bytes) - - -if __name__ == "__main__": - test_la( - bs=1, length=8192, qo_heads=64, kv_heads=64, dim=128, digest=False, bench=False + 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 d3a9543..e503eda 100644 --- a/tests/test_loss.py +++ b/tests/test_loss.py @@ -5,10 +5,10 @@ import random +import pytest import torch from linghe.facade.loss import moe_z_loss, softmax_cross_entropy -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.utils.loss import ( triton_softmax_cross_entropy_forward, @@ -38,15 +38,25 @@ def torch_z_loss(logits, coef=1e-6): 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=4096, - N=157184, - coef=1.0, - grad_coef=1.0, - ignore_index=None, - fill=False, - inplace=False, - bench=False, + M, N, coef, grad_coef, ignore_index, fill, benchmark, inplace=True ): device = "cuda:0" dtype = torch.bfloat16 @@ -105,34 +115,40 @@ def test_triton_softmax_cross_entropy( output_check(loss_ref, loss, name="loss", atol=1e-4, rtol=1e-5) output_check(grad_ref, grad, name="grad", digest=10) - if bench: - benchmark_func( - torch_cross_entropy, - logits.requires_grad_(), - targets, - output_grad, - ref_bytes=M * N * 2, - ) - benchmark_func( - triton_softmax_cross_entropy_forward, - logits, - targets, - ignore_index=ignore_index, - ref_bytes=M * N * 2, - ) - benchmark_func( - triton_softmax_cross_entropy_backward, - logits.detach().clone(), - targets, - sum_exp, - max_logit, - output_grad, - ignore_index=ignore_index, - ref_bytes=M * N * 4, - ) - - -def test_z_loss(L=4096, B=2, N=256, coef=0.001, bench=False): + 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 @@ -152,88 +168,12 @@ def test_z_loss(L=4096, B=2, N=256, coef=0.001, bench=False): output_check(loss_ref, loss, name="loss") output_check(grad_ref.float(), grad.float(), name="grad") - if bench: - benchmark_func(torch_z_loss, logits, coef=coef, ref_bytes=L * B * N * 4) - benchmark_func( - triton_moe_z_loss_forward, logits, coef=coef, ref_bytes=L * B * N * 4 - ) - benchmark_func( - triton_moe_z_loss_backward, - input_grad, - logits, - coef=coef, - ref_bytes=L * B * N * 8, - ) - - -if __name__ == "__main__": - test_triton_softmax_cross_entropy( - M=8192, N=157184, coef=1.0, grad_coef=1.0, inplace=True, bench=False - ) - test_triton_softmax_cross_entropy( - M=8192, N=157184, coef=1.0, grad_coef=1e-6, inplace=True, bench=False - ) - test_triton_softmax_cross_entropy( - M=8192, - N=157184, - coef=10000.0, - grad_coef=100.0, - fill=True, - inplace=True, - bench=False, - ) - test_triton_softmax_cross_entropy( - M=8192, - N=157184, - coef=1.0, - grad_coef=1.0, - fill=True, - ignore_index=-100, - inplace=True, - bench=False, - ) - test_triton_softmax_cross_entropy( - M=8192, - N=157184, - coef=1.0, - grad_coef=1.0, - fill=True, - ignore_index=0, - inplace=True, - bench=False, - ) - test_triton_softmax_cross_entropy( - M=8192, - N=157184 - 16, - coef=10000.0, - grad_coef=100.0, - fill=True, - inplace=True, - bench=False, + 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, ) - test_triton_softmax_cross_entropy( - M=8192, N=175175, coef=1.0, grad_coef=1.0, inplace=True, bench=False - ) - test_triton_softmax_cross_entropy( - M=8192, N=157184, coef=0.0, grad_coef=0.0, inplace=True, bench=False - ) - test_triton_softmax_cross_entropy( - M=8192, N=157184, coef=0.0, grad_coef=100.0, inplace=True, bench=False - ) - test_triton_softmax_cross_entropy( - M=8192, N=157184, coef=1000.0, grad_coef=0.0, inplace=True, bench=False - ) - test_triton_softmax_cross_entropy( - M=8192, - N=157184, - coef=100.0, - grad_coef=100.0, - fill=True, - inplace=True, - bench=False, - ) - test_triton_softmax_cross_entropy( - M=4096, N=157184, coef=0.1, grad_coef=1.0, inplace=True, bench=False - ) - - test_z_loss(L=4096, B=2, N=256, coef=1e-6, bench=False) diff --git a/tests/test_mla.py b/tests/test_mla.py index e234b4a..dd235b5 100644 --- a/tests/test_mla.py +++ b/tests/test_mla.py @@ -5,6 +5,7 @@ import math +import pytest import torch from linghe.attn.mla import ( @@ -15,7 +16,6 @@ triton_varlen_mla_backward, ) from linghe.facade.mla import multi_latend_attention -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check @@ -124,7 +124,13 @@ def head_wise_quant(x): return x_q, x_s -def test_softmax(M=128, N=128): +@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) @@ -134,7 +140,13 @@ def test_softmax(M=128, N=128): output_check(grad_ref, grad, atol=10, name="grad") -def test_dot_sum(M=128, N=128, D=128): +@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) @@ -143,17 +155,20 @@ def test_dot_sum(M=128, N=128, D=128): output_check(ds_ref, ds, atol=10, name="dot_sum") -def test_mla( - B=2, - L=4096, - H=16, - causal=True, - hpc=False, - safe=True, - coef=1.0, - clip_value=None, - bench=False, -): +@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 = ( @@ -200,48 +215,51 @@ def test_mla( 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") - if bench: - ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) - benchmark_func( - 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_func( - 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, - ) - - -def test_varlen_mla( - LS=[2048, 4096], - H=16, - causal=True, - hpc=False, - safe=True, - coef=1.0, - clip_value=None, - pad=False, - bench=False, -): + 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: @@ -319,7 +337,7 @@ def test_varlen_mla( max_q_length=max_q_length, causal=causal, safe=safe, - clip_value=clip_value, + clip_value=0.0 if clip_value is None else clip_value, ) output.backward(g) gq = q.grad @@ -334,47 +352,50 @@ def test_varlen_mla( 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") - if bench: - ref_flops = sum([L * L * H * (192 + 128) * (1 if causal else 2) for L in LS]) - benchmark_func( - 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_func( - 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, - ) - - -def test_fp8_mla( - B=2, L=4096, H=16, causal=True, hpc=False, quant_value=False, bench=False -): + 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) @@ -399,200 +420,15 @@ def test_fp8_mla( 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") - if bench: - ref_flops = B * L * L * H * (192 + 128) * (1 if causal else 2) - benchmark_func( - 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, - ) - - -if __name__ == "__main__": - test_softmax(M=128, N=128) - - test_dot_sum(M=128, N=128, D=128) - - test_mla( - B=1, - L=8192, - H=64, - causal=True, - hpc=False, - safe=False, - coef=1.0, - clip_value=500, - bench=False, - ) - test_mla( - B=1, - L=8192, - H=64, - causal=True, - hpc=False, - safe=False, - coef=1.0, - clip_value=500.0, - bench=False, - ) - test_mla( - B=1, - L=8192, - H=64, - causal=True, - hpc=True, - safe=False, - coef=1.0, - clip_value=None, - bench=False, - ) - test_mla( - B=1, - L=8192, - H=64, - causal=True, - hpc=False, - safe=True, - coef=100.0, - clip_value=None, - bench=False, - ) - test_mla( - B=1, - L=4096, - H=64, - causal=True, - hpc=False, - safe=False, - coef=1.0, - clip_value=None, - bench=False, - ) - test_mla( - B=1, - L=4096, - H=64, - causal=False, - hpc=False, - safe=False, - coef=1.0, - clip_value=None, - bench=False, - ) - test_mla( - B=1, - L=8192, - H=64, - causal=False, - hpc=False, - safe=False, - coef=1.0, - clip_value=None, - bench=False, - ) - test_mla( - B=1, - L=8192, - H=1, - causal=False, - hpc=False, - safe=False, - coef=1.0, - clip_value=None, - bench=False, - ) - - test_varlen_mla( - LS=[8192], - H=64, - causal=True, - hpc=False, - safe=False, - coef=1.0, - clip_value=None, - pad=False, - bench=False, - ) - test_varlen_mla( - LS=[8192], - H=64, - causal=True, - hpc=False, - safe=True, - coef=1.0, - clip_value=100.0, - pad=False, - bench=False, - ) - test_varlen_mla( - LS=[8192], - H=64, - causal=True, - hpc=False, - safe=True, - coef=1.0, - clip_value=None, - pad=True, - bench=False, - ) - - test_varlen_mla( - LS=[4096, 4096], H=64, causal=True, hpc=False, safe=True, coef=1.0, bench=False - ) - test_varlen_mla( - LS=[2048, 2048, 4096], - H=64, - causal=True, - hpc=True, - safe=True, - coef=1.0, - bench=False, - ) - test_varlen_mla( - LS=[127, 873, 3096], - H=64, - causal=False, - hpc=False, - safe=False, - coef=1.0, - bench=False, - ) - test_varlen_mla( - LS=[127, 873, 3456], - H=16, - causal=False, - hpc=False, - safe=True, - coef=1.0, - clip_value=100.0, - bench=False, - ) - test_varlen_mla( - LS=[127, 873, 3456], - H=16, - causal=False, - hpc=False, - safe=True, - coef=1.0, - pad=True, - bench=False, - ) - test_varlen_mla( - LS=[127, 873, 3456], - H=1, - causal=True, - hpc=False, - safe=True, - coef=1.0, - bench=False, - ) - - test_fp8_mla( - B=1, L=8192, H=64, causal=True, hpc=False, quant_value=False, bench=False + 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 index 0a10342..670a5b1 100644 --- a/tests/test_mul.py +++ b/tests/test_mul.py @@ -5,9 +5,9 @@ import random +import pytest import torch -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.utils.mul import triton_dot, triton_batch_scale, triton_inplace_scale @@ -25,7 +25,13 @@ def torch_batch_scale(xs, scale): return [x * scale for x in xs] -def test_dot(M=4096, N=4096, bench=False): +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + ], +) +def test_dot(M, N, benchmark): dtype = torch.bfloat16 device = "cuda:0" @@ -45,14 +51,17 @@ def test_dot(M=4096, N=4096, bench=False): 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) - ref_time = benchmark_func( - triton_dot, x, q, n_repeat=n_repeat, ref_time=ref_time - ) + 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) -def test_inplace_scale(M=2**20, bench=False): +@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) @@ -61,14 +70,18 @@ def test_inplace_scale(M=2**20, bench=False): ref_bytes = M * 8 - if bench: - ref_time = benchmark_func(torch_inplace_scale, x, scale, ref_bytes=ref_bytes) - benchmark_func( - triton_inplace_scale, x, scale, ref_bytes=ref_bytes, ref_time=ref_time - ) + 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) -def test_batch_scale(M=4096, N=2048, k=128, scale=1.0, bench=False): +@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( @@ -79,8 +92,6 @@ def test_batch_scale(M=4096, N=2048, k=128, scale=1.0, bench=False): ) for i in range(k) ] - # xs.append(torch.randn(2**32//N, N, - # dtype=dtype, device='cuda:0')) xs1 = [x.clone().detach() for x in xs] xs2 = [x.clone().detach() for x in xs] if scale == 0.0: @@ -99,15 +110,5 @@ def test_batch_scale(M=4096, N=2048, k=128, scale=1.0, bench=False): ref_bytes = sum([x.numel() for x in xs]) * 8 - if bench: - ref_time = benchmark_func(torch_batch_scale, xs, scale, ref_bytes=ref_bytes) - benchmark_func( - triton_batch_scale, xs, scale, ref_bytes=ref_bytes, ref_time=ref_time - ) - - -if __name__ == "__main__": - test_dot(M=4096, N=4096, bench=False) - test_inplace_scale(M=2**28 + 1, bench=False) - test_batch_scale(M=2048, N=1024, k=128, scale=2.0, bench=False) - test_batch_scale(M=2048, N=1024, k=128, scale=0.0, bench=False) + 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_rearange.py b/tests/test_rearange.py index e633038..c063fbd 100644 --- a/tests/test_rearange.py +++ b/tests/test_rearange.py @@ -3,9 +3,9 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.utils.rearange import triton_sort_chunks_by_index @@ -22,7 +22,13 @@ def torch_sort_chunks_by_index(x, scales, counts, indices): return output_data, output_scale -def test_sort_chunks_by_index(M=4096, N=4096, bench=False): +@pytest.mark.parametrize( + "M,N", + [ + (4096, 4096), + ], +) +def test_sort_chunks_by_index(M, N, benchmark): dtype = torch.bfloat16 device = "cuda:0" n_repeat = 100 @@ -49,30 +55,23 @@ def test_sort_chunks_by_index(M=4096, N=4096, bench=False): output_check(data_ref.view(torch.float8_e4m3fn), data, name="data") output_check(scale_ref, scale, name="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_sort_chunks_by_index, - x_q.view(torch.float8_e4m3fn), - x_scales, - split_size_list, - sorted_indices_list, - n_repeat=n_repeat, - ) - benchmark_func( - triton_sort_chunks_by_index, - x_q, - counts, - indices, - scales=x_scales, - n_repeat=n_repeat, - ) - - -if __name__ == "__main__": - test_sort_chunks_by_index(M=4096, N=4096) + 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 ae9e1e6..37d7947 100644 --- a/tests/test_reduce.py +++ b/tests/test_reduce.py @@ -5,9 +5,9 @@ import random +import pytest import torch -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.utils.reduce import ( triton_abs_max, @@ -36,7 +36,13 @@ def torch_count_zero(xs): return count -def test_triton_abs_max(M=4096, N=4096, bench=False): +@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)) ) @@ -45,11 +51,16 @@ def test_triton_abs_max(M=4096, N=4096, bench=False): 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) + benchmark(triton_abs_max, x, n_repeat=100, ref_bytes=M * N * 2) -def test_count_zero(M=4096, N=8192, k=32, bench=False): +@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) @@ -64,21 +75,24 @@ def test_count_zero(M=4096, N=8192, k=32, bench=False): # print(f'{count_ref=} {count=}') assert count_ref.item() - count.item() == 0 - 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_norm(M=4096, N=8192, coef=1.0, bench=False): + 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) @@ -89,25 +103,32 @@ def test_norm(M=4096, N=8192, coef=1.0, bench=False): sums = triton_norm(x, ord=1, norm=True, scalar=True) output_check(sum_ref, sums, "l1_norm") - if bench: - ref_bytes = M * N * 4 - n_repeat = 100 - ref_time = benchmark_func( - lambda x: x.norm(p=2), x, n_repeat=n_repeat, ref_bytes=ref_bytes - ) - benchmark_func( - triton_norm, - x, - ord=2, - norm=True, - scalar=True, - n_repeat=n_repeat, - ref_bytes=ref_bytes, - ref_time=ref_time, - ) - - -def test_batch_norm(M=4096, N=8192, k=32, coef=1.0, bench=False): + 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 @@ -126,23 +147,9 @@ def test_batch_norm(M=4096, N=8192, k=32, coef=1.0, bench=False): sums = triton_batch_norm(xs, ord=-1, norm=False) output_check(sum_ref, sums, "inf_norm") - if bench: - ref_bytes = sum([x.numel() for x in xs]) * 4 - n_repeat = 100 - ref_time = benchmark_func(torch_sum, xs, 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__": - test_triton_abs_max(M=4096, N=4096, bench=False) - test_count_zero(M=4096, N=8192, k=32, bench=False) - test_norm(M=100000, N=8192, bench=False) - test_batch_norm(M=4096, N=1024, k=16, bench=False) - test_batch_norm(M=4096, N=1024, k=64, bench=False) - test_batch_norm(M=4096, N=2048, k=1024, coef=1e12, bench=False) + ref_bytes = sum([x.numel() for x in xs]) * 4 + 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 6f9f127..21852ec 100644 --- a/tests/test_rope.py +++ b/tests/test_rope.py @@ -3,10 +3,10 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch from linghe.facade.rope import qk_norm_half_rope -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.utils.rope import ( triton_half_rope_forward, @@ -299,9 +299,14 @@ def torch_varlen_qk_norm_and_half_rope( return qoss, koss, voss -def test_half_rope( - B=2, L=4096, H=32, h=8, D=128, rope_theta=10000.0, transposed=True, bench=False -): +@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) @@ -330,29 +335,32 @@ def test_half_rope( 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, + ) + + +@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=2, - L=4096, - H=32, - h=8, - D=128, - rope_theta=10000.0, - eps=1e-6, - interleaved=True, - transposed=True, - silu=False, - bench=False, + B, L, H, h, D, rope_theta, eps, interleaved, transposed, silu, benchmark ): dtype = torch.bfloat16 device = "cuda:0" @@ -442,53 +450,55 @@ def test_qk_norm_and_half_rope( output_check(dqw_ref, dqw, name="dqw") output_check(dkw_ref, dkw, name="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, - silu=silu, - 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, - transposed=transposed, - interleaved=interleaved, - silu=silu, - ref_bytes=L * B * (H + 2 * h) * D * 6, - n_profile=0, - ) + 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=[2048, 2048], - H=32, - h=4, - dim=128, - rope_theta=10000.0, - silu=False, - interleaved=True, - bench=False, - cp_size=1, - cp_rank=0, + lengths, H, h, dim, rope_theta, silu, interleaved, cp_size, cp_rank, benchmark ): - dtype = torch.bfloat16 + # 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) @@ -567,51 +577,57 @@ def test_varlen_qk_norm_and_half_rope( 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=5.0, rtol=0.02) - output_check(dkw_ref, dkw.to(dtype), name="dkw", atol=5.0, rtol=0.02) - - if bench: - lbh = sum(lengths) // cp_size * H - benchmark_func( - 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_func( - 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, - ) + 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, + ) -def test_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=False, bench=False): +@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) @@ -665,32 +681,42 @@ def test_mla_rope(B=2, L=4096, H=32, rope_theta=10000.0, transpose=False, bench= output_check(dkv_ref, dkv, name="dkv") output_check(dp_ref, dp, name="dp") - if bench: - lbh = L * B * H - benchmark_func( - 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_func( - 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, - ) - - -def test_varlen_mla_rope( - lengths=[2048, 2048], H=32, rope_theta=10000.0, bench=False, cp_size=1, cp_rank=0 -): - dtype = torch.bfloat16 + 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 @@ -776,245 +802,34 @@ def test_varlen_mla_rope( 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) - if bench: - lbh = sum(lengths) // cp_size * H - benchmark_func( - 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_func( - 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, - ) - - -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=2, - L=4096, - H=16, - h=16, - D=128, - rope_theta=10000.0, - interleaved=True, - transposed=True, - silu=True, - bench=False, - ) - test_qk_norm_and_half_rope( - B=2, - L=4096, - H=16, - h=16, - D=128, - rope_theta=10000.0, - interleaved=True, - transposed=True, - silu=False, - 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=False, - silu=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=False, - silu=False, - bench=False, - ) - test_qk_norm_and_half_rope( - B=4, - L=4096, - H=32, - h=4, - D=128, - rope_theta=10000.0, - interleaved=False, - transposed=True, - silu=True, - bench=False, - ) - test_qk_norm_and_half_rope( - B=4, - L=4096, - H=24, - h=6, - D=128, - rope_theta=10000.0, - interleaved=True, - transposed=True, - silu=False, - bench=False, - ) - test_qk_norm_and_half_rope( - B=4, - L=4096, - H=32, - h=32, - D=128, - rope_theta=10000.0, - interleaved=False, - transposed=False, - silu=True, - bench=False, + 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, ) - test_qk_norm_and_half_rope( - B=1, - L=4096, - H=32, - h=32, - D=128, - rope_theta=10000.0, - interleaved=False, + 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, - silu=False, - bench=False, - ) - test_varlen_qk_norm_and_half_rope( - lengths=[2048], - H=24, - h=6, - dim=128, - rope_theta=10000.0, - silu=False, - interleaved=True, - cp_size=1, - cp_rank=0, - bench=False, - ) - test_varlen_qk_norm_and_half_rope( - lengths=[1024, 4096, 4096, 568], - H=32, - h=4, - dim=128, - rope_theta=10000.0, - silu=False, - interleaved=True, - cp_size=1, - cp_rank=0, - bench=False, - ) - test_varlen_qk_norm_and_half_rope( - lengths=[2048, 4096, 4096], - H=32, - h=4, - dim=128, - rope_theta=10000.0, - silu=False, - interleaved=True, - cp_size=1, - cp_rank=0, - bench=False, - ) - test_varlen_qk_norm_and_half_rope( - lengths=[2048, 3072, 4096], - H=32, - h=4, - dim=128, - rope_theta=10000.0, - silu=False, - interleaved=True, - cp_size=4, - cp_rank=0, - bench=False, - ) - test_varlen_qk_norm_and_half_rope( - lengths=[2048, 4096, 4096], - H=32, - h=4, - dim=128, - rope_theta=10000.0, - silu=True, - interleaved=False, - cp_size=4, - cp_rank=0, - bench=False, - ) - test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=False, bench=False) - test_mla_rope(B=4, L=4096, H=16, rope_theta=10000.0, transpose=True, bench=False) - test_varlen_mla_rope( - lengths=[8192], H=64, rope_theta=10000.0, cp_size=1, cp_rank=0, bench=False - ) - test_varlen_mla_rope( - lengths=[4096, 4096], - H=16, - rope_theta=10000.0, - cp_size=1, - cp_rank=0, - bench=False, - ) - test_varlen_mla_rope( - lengths=[4096 * 4, 2048 * 4, 2048 * 4], - H=32, - rope_theta=10000.0, - cp_size=4, - cp_rank=0, - bench=False, - ) - test_varlen_mla_rope( - lengths=[4096 * 4, 2048 * 4, 2048 * 4], - H=32, - rope_theta=10000.0, - cp_size=4, - cp_rank=1, - bench=False, - ) - test_varlen_mla_rope( - lengths=[4096 * 4, 2048 * 4, 2048 * 4], - H=32, - rope_theta=10000.0, - cp_size=4, - cp_rank=2, - bench=False, - ) - test_varlen_mla_rope( - lengths=[4096 * 4, 2048 * 4, 2048 * 4], - H=32, - rope_theta=10000.0, - cp_size=4, - cp_rank=3, - bench=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 9c71505..a518274 100644 --- a/tests/test_scatter.py +++ b/tests/test_scatter.py @@ -3,13 +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.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" @@ -24,7 +25,15 @@ def torch_scatter_add(x, outputs, indices, weights): return outputs.to(dtype) -def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): +@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" @@ -49,22 +58,13 @@ def test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False): output_check(sums_ref, sums_unpermute, "unpermute_data") output_check(probs, output_prob, "unpermute_prob") - if bench: - 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, - ) - - -if __name__ == "__main__": - test_scatter(M=4098, N=4096, n_experts=32, topk=2, bias=0.0, bench=False) - test_scatter(M=2467, N=4096, n_experts=32, topk=2, bias=-0.1, bench=False) - test_scatter(M=2467, N=1536, n_experts=32, topk=2, bias=-0.1, bench=False) + 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 00ac2af..ab9aa48 100644 --- a/tests/test_silu.py +++ b/tests/test_silu.py @@ -7,9 +7,9 @@ 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, @@ -17,66 +17,72 @@ 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 +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): +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 + 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.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_q, y_scale = 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 + return y_q, y_scale -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 +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) @@ -84,24 +90,32 @@ def torch_silu_and_block_quant_forward(x, round_scale=True): return y_q, y_scale, yt_q, yt_scale +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( + q, dx_scale = torch_smooth_quant( dx, smooth_scale, reverse=reverse, round_scale=round_scale ) - yt_q, yt_scale, ms = torch_smooth_quant( + yt_q, yt_scale = torch_smooth_quant( dx.t().contiguous(), transpose_smooth_scale, reverse=reverse, @@ -110,10 +124,10 @@ def torch_silu_and_smooth_quant_backward( 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 @@ -123,8 +137,19 @@ def torch_silu_and_block_quant_backward(grad, x, round_scale=True): return q, dx_scale, yt_q, yt_scale +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, round_scale=True, reverse=False + xs, weight, counts, smooth_scales=None, limit=None, round_scale=True, reverse=False ): counts = counts.tolist() N = xs.shape[1] @@ -132,8 +157,7 @@ def torch_batch_weighted_silu_and_smooth_quant_forward( 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() @@ -141,27 +165,24 @@ def torch_batch_weighted_silu_and_smooth_quant_forward( 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 = 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 + xs, weight, counts, limit=None, round_scale=True ): counts = counts.tolist() N = xs.shape[1] @@ -183,7 +204,7 @@ def torch_batch_weighted_silu_and_block_quant_forward( s = 0 for i, c in enumerate(counts): x = xs[s : s + c] - y = torch_weighted_silu(x, weight[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) @@ -199,6 +220,43 @@ def torch_batch_weighted_silu_and_block_quant_forward( return qs, scales, qts, qtscales +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, @@ -206,6 +264,7 @@ def torch_batch_weighted_silu_and_smooth_quant_backward( counts, smooth_scales=None, transpose_smooth_scale=None, + limit=None, round_scale=True, reverse=False, ): @@ -225,14 +284,14 @@ def torch_batch_weighted_silu_and_smooth_quant_backward( 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( + 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() @@ -241,7 +300,7 @@ def torch_batch_weighted_silu_and_smooth_quant_backward( 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( + qt, t_scale = torch_smooth_quant( dxt, dxt_s, reverse=reverse, round_scale=round_scale ) @@ -258,7 +317,7 @@ def torch_batch_weighted_silu_and_smooth_quant_backward( def torch_batch_weighted_silu_and_block_quant_backward( - grad_output, x, weight, counts, round_scale=True + grad_output, x, weight, counts, limit=None, round_scale=True ): if sum(counts) == 0: device = x.device @@ -274,7 +333,7 @@ def torch_batch_weighted_silu_and_block_quant_backward( 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 = [] @@ -296,11 +355,59 @@ def torch_batch_weighted_silu_and_block_quant_backward( return dx_q, dx_scale, dw, qts, qtscales -def test_weighted_silu(M=4096, N=4096, asm=False, coef=1.0, bench=False): - x = torch.randn((M, N), 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 // 2), dtype=torch.bfloat16, 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, asm=asm) output_check(ref_y, y, "y") @@ -310,47 +417,52 @@ def test_weighted_silu(M=4096, N=4096, asm=False, coef=1.0, bench=False): output_check(dx_ref, dx, "dx") output_check(dw_ref, dw, "dw", rtol=3e-3, atol=3e-3) - if bench: - benchmark_func( - triton_weighted_silu_forward, - x, - weight, - asm=asm, - 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, - ) + 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, + ) -def test_silu_and_smooth_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") +@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 // 2), dtype=torch.bfloat16, device="cuda:0") * grad_coef - ) - 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") + 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, y_maxs_ref = torch_silu_and_smooth_quant_forward( + 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, y_maxs = triton_silu_and_smooth_quant_forward( - x, smooth_scale=smooth_scale, round_scale=round_scale, calibrate=True + 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") - 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( @@ -376,58 +488,68 @@ def test_silu_and_smooth_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, bench=Fa output_check(dxt_q_ref, dxt_q, "smooth.dxt_data", rtol=0.125) 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, coef=1.0, grad_coef=1.0, bench=False): - x = torch.randn((M, N), dtype=torch.bfloat16, device="cuda:0") - x = (x * coef).clone().detach().requires_grad_() - grad_output = ( - torch.randn((M, N // 2), dtype=torch.bfloat16, device="cuda:0") * grad_coef + 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=round_scale + 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, output_mode=0 + 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, output_mode=1 + 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, output_mode=2 + 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") @@ -435,52 +557,126 @@ def test_silu_and_block_quant(M=4096, N=4096, coef=1.0, grad_coef=1.0, bench=Fal 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) + 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=round_scale + 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") - if bench: - benchmark_func( - triton_silu_and_block_quant_forward, - x, - round_scale=round_scale, - output_mode=0, - n_repeat=100, - ref_bytes=M * N * 3, - ) - benchmark_func( - triton_silu_and_block_quant_forward, - x, - round_scale=round_scale, - output_mode=1, - n_repeat=100, - ref_bytes=M * N * 3, - ) - benchmark_func( - triton_silu_and_block_quant_forward, - x, - round_scale=round_scale, - output_mode=2, - 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, - ) + 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=1024, N=4096, n_experts=32, coef=1.0, grad_coef=1.0, bench=False + 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) @@ -488,34 +684,32 @@ def test_triton_batch_weighted_silu_and_smooth_quant( 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") * coef + 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 // 2), dtype=torch.float32, device="cuda:0") * 10 + 1 + torch.rand((n_experts, N), dtype=torch.float32, device="cuda:0") * 10 ) grad_output = ( - torch.randn((bs, N // 2), dtype=torch.bfloat16, device="cuda:0") * grad_coef + torch.randn((bs, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef ) grad_smooth_scales = ( - 1 + torch.rand((n_experts, N), dtype=torch.float32, device="cuda:0") * 10 + 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 rtol = 2 if round_scale else 0.125 - x_q_ref, x_scale_ref, x_max_ref = ( - torch_batch_weighted_silu_and_smooth_quant_forward( - x, - weight, - counts, - smooth_scales=smooth_scales, - round_scale=round_scale, - reverse=False, - ) + 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_q, x_scale = triton_batch_weighted_silu_and_smooth_quant_forward( x, weight, counts, @@ -558,35 +752,44 @@ def test_triton_batch_weighted_silu_and_smooth_quant( output_check(dxt_ref, dxt, "smooth.dxt", rtol=rtol) 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, - ) + 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=1024, N=4096, n_experts=32, bench=False, coef=1.0, grad_coef=1.0 + 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) @@ -594,51 +797,75 @@ def test_triton_batch_weighted_silu_and_block_quant( 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") * coef + 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 // 2), dtype=torch.bfloat16, device="cuda:0") * grad_coef + 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, round_scale=round_scale + 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, round_scale=round_scale, output_mode=2 + x, + weight, + counts, + count_list, + limit=limit, + round_scale=round_scale, + 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") + output_check(xt_scale_ref, xt_scale.view(-1), "block.t_scale", atol=1e-3) x_q, x_scale, xt_q, xt_scale = triton_batch_weighted_silu_and_block_quant_forward( - x, weight, counts, count_list, round_scale=round_scale, output_mode=0 + x, + weight, + counts, + 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, round_scale=round_scale, output_mode=1 + x, + weight, + counts, + count_list, + limit=limit, + round_scale=round_scale, + 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 + 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, round_scale=round_scale + grad_output, + x, + weight, + counts, + splits=count_list, + limit=limit, + round_scale=round_scale, ) ) output_check(dx_ref, dx, "block.dx", rtol=rtol) @@ -648,94 +875,149 @@ def test_triton_batch_weighted_silu_and_block_quant( output_check(dxt_ref, dxt.view(-1), "block.dxt", rtol=rtol) output_check(dxt_scale_ref, dxt_scale.view(-1), "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=1, - n_repeat=100, - ref_bytes=n_experts * M * N * 3, - 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=4096, coef=1.0, asm=False, bench=False) - test_weighted_silu(M=16384, N=4096, coef=1.0, asm=True, bench=False) - test_weighted_silu(M=8192, N=1536, bench=False) - test_weighted_silu(M=0, N=1536, bench=False) + 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) - 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) + x = torch.randn((bs, 2 * N), dtype=torch.bfloat16, device="cuda:0") * coef + weight = torch.randn((bs, 1), dtype=torch.float32, device="cuda:0") - test_silu_and_block_quant(M=16384, N=1024, bench=False) - test_silu_and_block_quant(M=8192, N=4096, bench=False) - test_silu_and_block_quant(M=16384, N=1536, bench=False) - test_silu_and_block_quant(M=4096, N=1536 * 8, bench=False) - test_silu_and_block_quant( - M=4096, N=1536 * 8, coef=100.0, grad_coef=100.0, bench=False + grad_output = ( + torch.randn((bs, N), dtype=torch.bfloat16, device="cuda:0") * grad_coef ) - test_silu_and_block_quant(M=4096, N=1536 * 8, coef=0.0, grad_coef=0.0, bench=False) - test_triton_batch_weighted_silu_and_smooth_quant( - M=0, N=2048, n_experts=32, bench=False + 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 + ) ) - test_triton_batch_weighted_silu_and_smooth_quant( - M=2048, N=2048, n_experts=32, bench=False + 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 ) - test_triton_batch_weighted_silu_and_block_quant( - M=0, N=1536, n_experts=32, bench=False - ) - test_triton_batch_weighted_silu_and_block_quant( - M=2048, N=8192, n_experts=32, bench=False + 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 + ) ) - test_triton_batch_weighted_silu_and_block_quant( - M=12080, N=1536, n_experts=32, coef=100.0, grad_coef=100.0, bench=False + 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 + ) ) - test_triton_batch_weighted_silu_and_block_quant( - M=12080, N=1536, n_experts=32, coef=0.0, grad_coef=0.0, bench=False + 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 a16c826..563c2df 100644 --- a/tests/test_smooth_quant.py +++ b/tests/test_smooth_quant.py @@ -3,17 +3,18 @@ Copyright (c) Ant Financial Service Group and its affiliates. """ +import pytest import torch -from linghe.facade.smooth_quant_linear import SmoothQuantLinear +from linghe.facade.linear import SmoothQuantLinear 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, + triton_batch_smooth_quant, + triton_batch_transpose_smooth_quant, + triton_transpose_rescale_smooth_quant, + triton_subrow_smooth_quant, ) -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check from linghe.tools.util import torch_make_indices, torch_smooth_quant, round_up @@ -32,7 +33,41 @@ 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( @@ -95,7 +130,7 @@ def torch_rescale_quant( ): assert reverse y = y_q.float() / org_smooth_scale * y_scale[:, None] - y_q, y_scale, _ = torch_smooth_quant( + y_q, y_scale = torch_smooth_quant( y.t(), transpose_smooth_scale, reverse=True, round_scale=round_scale ) return y_q, y_scale @@ -105,42 +140,59 @@ 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): +@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() + 1.0 round_scale = False rtol = 2 if round_scale else 0.125 - x_q_ref, scales_ref, x_maxs_ref = torch_smooth_quant( + x_q_ref, scales_ref = torch_smooth_quant( x, smooth_scale, reverse=False, round_scale=round_scale ) - x_q, x_scale, x_maxs = triton_smooth_quant( - x, smooth_scale, reverse=False, round_scale=round_scale, calibrate=True + 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") - 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, - ) + + benchmark( + triton_smooth_quant, + x, + smooth_scale, + reverse=False, + round_scale=True, + ref_bytes=M * N * 3, + ) -def test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, size=16384): +@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( @@ -200,7 +252,16 @@ def test_triton_subrow_smooth_quant(M=4096, N=5120, offset=4096, size=16384): 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): +@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 @@ -210,7 +271,7 @@ def test_triton_transpose_smooth_quant(M=4096, N=4096, bench=False): 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( + q_ref, scale_ref = torch_smooth_quant( y.T.contiguous(), transpose_smooth_scale, reverse=True, round_scale=True ) @@ -220,19 +281,27 @@ def test_triton_transpose_smooth_quant(M=4096, N=4096, bench=False): 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, - ) + benchmark( + 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): +@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 @@ -249,11 +318,11 @@ def test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, round_scale=False torch.ceil(torch.log2(transpose_smooth_scale)) ) - y_q, y_scale, y_maxs = triton_smooth_quant( + y_q, y_scale = triton_smooth_quant( y, org_smooth_scale, reverse=True, round_scale=round_scale ) - yt_gt, yt_scale_gt, yt_maxs_gt = torch_smooth_quant( + yt_gt, yt_scale_gt = torch_smooth_quant( y.t(), transpose_smooth_scale, reverse=True, round_scale=round_scale ) @@ -288,11 +357,27 @@ def test_triton_transpose_rescale_smooth_quant(M=4096, N=4096, round_scale=False # 'triton_transpose_rescale_smooth_quant.data.gt') # output_check(yt_scale_gt, yt_scale, # 'triton_transpose_rescale_smooth_quant.scale.gt') + 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, + ) -def test_triton_batch_smooth_quant( - M=4096, N=4096, n_experts=32, topk=8, round_scale=False, bench=False -): +@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( @@ -308,57 +393,109 @@ def test_triton_batch_smooth_quant( (sum(token_count_per_expert_list), N), dtype=torch.bfloat16, device=device ) - x_q, x_scale, x_maxs = triton_batch_smooth_quant( + 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 = 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) + 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, - calibrate=True, + 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, ) 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 + 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 ) - x_scale_ref = torch.cat(x_scale_ref, 0) 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_q_ref, x_q, "batch_transpose_smooth_quant.data", rtol=rtol) output_check( - x_scale_ref.float(), x_scale.float(), "triton_batch_smooth_quant.scale" + x_scale_ref.float(), x_scale.float(), "batch_transpose_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, - ) + 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, + ) -def test_smooth_quant_linear(M=8192, N=1024, K=2048): +@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) @@ -369,41 +506,12 @@ def test_smooth_quant_linear(M=8192, N=1024, K=2048): y_ref = x @ w.t() y = linear(x) - output_check(y_ref, y, name="y") + output_check(y_ref, y, name="y", atol=-1) dx_ref = dy @ w dw_ref = dy.t() @ x y.backward(dy) dw = linear.weight.grad dx = x.grad - output_check(dx_ref, dx, name="dx") - output_check(dw_ref, dw, name="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) + output_check(dx_ref, dx, name="dx", atol=-1) + output_check(dw_ref, dw, name="dw", atol=-1) diff --git a/tests/test_unary.py b/tests/test_unary.py index 7065220..3d142fc 100644 --- a/tests/test_unary.py +++ b/tests/test_unary.py @@ -5,103 +5,107 @@ import random +import pytest import torch -from linghe.tools.benchmark import benchmark_func from linghe.tools.check import output_check -from linghe.utils.unary import triton_calculate_smooth_scale, triton_batch_clip +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, round_scale=False): +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) + 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))) + weight_smooth_scales = torch.exp2( + torch.ceil(torch.log2(weight_smooth_scales))) return weight_smooth_scales +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 -def test_calculate_smooth_scale(N=4096, bench=False): - 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.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") + 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, - ) - - -def test_batch_clip(M=2048, N=1024, k=1024, clip_value=1.0, inf=False, bench=False): - 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) - ] + 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") + 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 bench: - 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_func( - 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_func( - triton_batch_clip, - xs4, - clip_value, - ref_bytes=ref_bytes, - ref_time=ref_time, - n_repeat=n_repeat, - n_warmup=0, - ) - - -if __name__ == "__main__": - # test_calculate_smooth_scale(N=4096*32) - # test_calculate_smooth_scale(N=4096*32-1897) - # test_batch_clip(M=2048, N=8192, k=128, clip_value=0.1, bench=False) - # test_batch_clip(M=2048, N=1024, k=128, clip_value=1.0, bench=False) - test_batch_clip(M=2048, N=1024, k=128, clip_value=100.0, inf=True, bench=False) + 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 + 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) From eaae73f15bd51d873027c666fbd33a1712f46e8c Mon Sep 17 00:00:00 2001 From: "liangchen.liangche" Date: Mon, 27 Apr 2026 17:48:07 +0800 Subject: [PATCH 10/11] fix gradient --- linghe/facade/gate.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/linghe/facade/gate.py b/linghe/facade/gate.py index 47cbaed..6bfcf34 100644 --- a/linghe/facade/gate.py +++ b/linghe/facade/gate.py @@ -89,14 +89,14 @@ def forward( output = cls( shape=x_q.shape, - dtype=input.dtype, + 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=input.requires_grad, + requires_grad=attn_output.requires_grad, ) ctx.save_for_backward(attn_output, gate, weight) @@ -131,4 +131,4 @@ def backward(ctx, dy): requires_grad=ctx.input_requires_grad, ) - return dx, dg_out, dw, None, None + return dx, dg_out, dw, None, None, None, None, None From 059a7096d257121415c39a6a3cd8ac63091a0754 Mon Sep 17 00:00:00 2001 From: "liangchen.liangche" Date: Tue, 28 Apr 2026 11:52:04 +0800 Subject: [PATCH 11/11] fix typo --- linghe/facade/gate.py | 2 +- linghe/facade/gemm.py | 2 -- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/linghe/facade/gate.py b/linghe/facade/gate.py index 6bfcf34..5843f1f 100644 --- a/linghe/facade/gate.py +++ b/linghe/facade/gate.py @@ -128,7 +128,7 @@ def backward(ctx, dy): 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=ctx.input_requires_grad, + 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 index 91fcd4d..85f982a 100644 --- a/linghe/facade/gemm.py +++ b/linghe/facade/gemm.py @@ -86,8 +86,6 @@ def smooth_groued_gemm( if m == 0: continue - x_q = B[i]._rowwise_data - x_scale = B[i]._rowwise_scale_inv x_q = B[i]._rowwise_data x_scale = B[i]._rowwise_scale_inv w_q = A[i]._rowwise_data