已开启
feat(triton): add bwd_dtrap_ddt_kernel from mamba to mindspeed ops #91
Tsuki创建于 7月2日
feat(triton): add bwd_dtrap_ddt_kernel from mamba to mindspeed ops #91
已开启
共 6 个文件变更+588-0
| @@ -0,0 +1,77 @@ | |||||||||
| 1 | +# Copyright (c) 2023, Tri Dao, Albert Gu | ||||||||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. | ||||||||
| 3 | +# | ||||||||
L | |||||||||
| 4 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||||||||
| 5 | +# you may not use this file except in compliance with the License. | ||||||||
| 6 | + | ||||||||
| 7 | +import torch | ||||||||
| 8 | + | ||||||||
| 9 | +from mindspeed_ops.arch32.triton.bwd_dtrap_ddt import bwd_dtrap_ddt_kernel | ||||||||
| 10 | +from mindspeed_ops.utils import is_arch35 | ||||||||
| 11 | + | ||||||||
| 12 | +__all__ = ["bwd_dtrap_ddt"] | ||||||||
| 13 | + | ||||||||
| 14 | + | ||||||||
| 15 | +def bwd_dtrap_ddt( | ||||||||
| 16 | + trap: torch.Tensor, | ||||||||
| 17 | + dt: torch.Tensor, | ||||||||
| 18 | + dfactor: torch.Tensor, | ||||||||
| 19 | + dgamma_diag: torch.Tensor, | ||||||||
| 20 | + ddt: torch.Tensor, | ||||||||
| 21 | + dtrap: torch.Tensor, | ||||||||
| 22 | + chunk_size: int, | ||||||||
🟡 Medium Priority API 函数 同类 API (如 建议:在 改动建议
![]() ![]() 不准确? | |||||||||
| 23 | +) -> tuple: | ||||||||
| 24 | + if is_arch35(): | ||||||||
| 25 | + raise NotImplementedError("bwd_dtrap_ddt is not supported on arch35") | ||||||||
| 26 | + | ||||||||
| 27 | + trap = trap.contiguous() | ||||||||
| 28 | + dt = dt.contiguous() | ||||||||
| 29 | + dfactor = dfactor.contiguous() | ||||||||
| 30 | + dgamma_diag = dgamma_diag.contiguous() | ||||||||
| 31 | + ddt = ddt.contiguous() | ||||||||
| 32 | + dtrap = dtrap.contiguous() | ||||||||
| 33 | + | ||||||||
| 34 | + B, H, S = trap.shape | ||||||||
| 35 | + | ||||||||
| 36 | + assert dt.shape == (B, H, S), f"dt shape mismatch: expected ({B}, {H}, {S}), got {dt.shape}" | ||||||||
| 37 | + assert dfactor.shape == (B, H, S), f"dfactor shape mismatch: expected ({B}, {H}, {S}), got {dfactor.shape}" | ||||||||
| 38 | + assert dgamma_diag.shape == (B, H, S), ( | ||||||||
| 39 | + f"dgamma_diag shape mismatch: expected ({B}, {H}, {S}), got {dgamma_diag.shape}" | ||||||||
| 40 | + ) | ||||||||
| 41 | + assert ddt.shape == (B, H, S), f"ddt shape mismatch: expected ({B}, {H}, {S}), got {ddt.shape}" | ||||||||
| 42 | + assert dtrap.shape == (B, H, S), f"dtrap shape mismatch: expected ({B}, {H}, {S}), got {dtrap.shape}" | ||||||||
| 43 | + | ||||||||
| 44 | + nchunks = (S + chunk_size - 1) // chunk_size | ||||||||
| 45 | + | ||||||||
| 46 | + grid = (B, H, nchunks) | ||||||||
| 47 | + | ||||||||
| 48 | + bwd_dtrap_ddt_kernel[grid]( | ||||||||
| 49 | + trap, | ||||||||
| 50 | + dt, | ||||||||
| 51 | + dfactor, | ||||||||
| 52 | + dgamma_diag, | ||||||||
| 53 | + ddt, | ||||||||
| 54 | + dtrap, | ||||||||
| 55 | + trap.stride(0), | ||||||||
| 56 | + trap.stride(1), | ||||||||
| 57 | + trap.stride(2), | ||||||||
| 58 | + dt.stride(0), | ||||||||
| 59 | + dt.stride(1), | ||||||||
| 60 | + dt.stride(2), | ||||||||
| 61 | + dfactor.stride(0), | ||||||||
| 62 | + dfactor.stride(1), | ||||||||
| 63 | + dfactor.stride(2), | ||||||||
| 64 | + dgamma_diag.stride(0), | ||||||||
| 65 | + dgamma_diag.stride(1), | ||||||||
| 66 | + dgamma_diag.stride(2), | ||||||||
| 67 | + ddt.stride(0), | ||||||||
| 68 | + ddt.stride(1), | ||||||||
| 69 | + ddt.stride(2), | ||||||||
| 70 | + dtrap.stride(0), | ||||||||
| 71 | + dtrap.stride(1), | ||||||||
| 72 | + dtrap.stride(2), | ||||||||
| 73 | + SEQLEN=S, | ||||||||
| 74 | + CHUNK_SIZE=chunk_size, | ||||||||
| 75 | + ) | ||||||||
| 76 | + | ||||||||
| 77 | + return ddt, dtrap | ||||||||
| @@ -0,0 +1,101 @@ | |||
| 1 | +# Copyright (c) 2023, Tri Dao, Albert Gu | ||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. | ||
| 3 | +# | ||
| 4 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | +# you may not use this file except in compliance with the License. | ||
| 6 | +# pylint: disable=import-error | ||
| 7 | +import triton | ||
| 8 | +import triton.language as tl # pylint: disable=E0611 | ||
| 9 | + | ||
| 10 | + | ||
| 11 | +# pylint: disable=E1101 | ||
| 12 | +def bwd_dtrap_ddt_kernel( | ||
【openlibing.ci】识别到代码检查告警抑制注释,匹配工具:pylint,请Committer检视其合理性。 ![]() ![]() | |||
| 13 | + trap_ptr, | ||
| 14 | + dt_ptr, | ||
| 15 | + dfactor_ptr, | ||
| 16 | + dgamma_diag_ptr, | ||
| 17 | + ddt_ptr, | ||
| 18 | + dtrap_ptr, | ||
| 19 | + stride_trap_batch, | ||
| 20 | + stride_trap_head, | ||
| 21 | + stride_trap_seq, | ||
| 22 | + stride_dt_batch, | ||
| 23 | + stride_dt_head, | ||
| 24 | + stride_dt_seq, | ||
| 25 | + stride_dfactor_batch, | ||
| 26 | + stride_dfactor_head, | ||
| 27 | + stride_dfactor_seq, | ||
| 28 | + stride_dgamma_diag_batch, | ||
| 29 | + stride_dgamma_diag_head, | ||
| 30 | + stride_dgamma_diag_seq, | ||
| 31 | + stride_ddt_batch, | ||
| 32 | + stride_ddt_head, | ||
| 33 | + stride_ddt_seq, | ||
| 34 | + stride_dtrap_batch, | ||
| 35 | + stride_dtrap_head, | ||
| 36 | + stride_dtrap_seq, | ||
| 37 | + SEQLEN: tl.constexpr, | ||
| 38 | + CHUNK_SIZE: tl.constexpr, | ||
| 39 | +): | ||
| 40 | + pid_batch = tl.program_id(0) | ||
| 41 | + pid_head = tl.program_id(1) | ||
| 42 | + pid_chunk = tl.program_id(2) | ||
| 43 | + | ||
| 44 | + chunk_start = pid_chunk * CHUNK_SIZE | ||
| 45 | + offs_c = tl.arange(0, CHUNK_SIZE) | ||
| 46 | + offs_seq = chunk_start + offs_c | ||
| 47 | + | ||
| 48 | + trap_offset = pid_batch * stride_trap_batch + pid_head * stride_trap_head | ||
| 49 | + dt_offset = pid_batch * stride_dt_batch + pid_head * stride_dt_head | ||
| 50 | + dfactor_offset = pid_batch * stride_dfactor_batch + pid_head * stride_dfactor_head | ||
| 51 | + dgamma_diag_offset = pid_batch * stride_dgamma_diag_batch + pid_head * stride_dgamma_diag_head | ||
| 52 | + | ||
| 53 | + strap_block = tl.load( | ||
| 54 | + trap_ptr + trap_offset + (offs_seq + 1) * stride_trap_seq, mask=(offs_seq + 1) < SEQLEN, other=0.0 | ||
| 55 | + ) | ||
| 56 | + sdt_block = tl.load(dt_ptr + dt_offset + (offs_seq + 1) * stride_dt_seq, mask=(offs_seq + 1) < SEQLEN, other=0.0) | ||
| 57 | + trap_block = tl.load(trap_ptr + trap_offset + offs_seq * stride_trap_seq, mask=offs_seq < SEQLEN, other=0.0) | ||
| 58 | + dt_block = tl.load(dt_ptr + dt_offset + offs_seq * stride_dt_seq, mask=offs_seq < SEQLEN, other=0.0) | ||
| 59 | + dfactor_block = tl.load( | ||
| 60 | + dfactor_ptr + dfactor_offset + offs_seq * stride_dfactor_seq, mask=offs_seq < SEQLEN, other=0.0 | ||
| 61 | + ) | ||
| 62 | + dgamma_diag_input_block = tl.load( | ||
| 63 | + dgamma_diag_ptr + dgamma_diag_offset + offs_seq * stride_dgamma_diag_seq, mask=offs_seq < SEQLEN, other=0.0 | ||
| 64 | + ) | ||
| 65 | + | ||
| 66 | + dgamma_block = dfactor_block + dgamma_diag_input_block | ||
| 67 | + dsgamma_block = dfactor_block | ||
| 68 | + | ||
| 69 | + dsdt_block = tl.sigmoid(-strap_block.to(tl.float32)) * dsgamma_block | ||
| 70 | + dstrap_block = -sdt_block * dsgamma_block | ||
| 71 | + | ||
| 72 | + prev_seq = chunk_start - 1 | ||
| 73 | + prev_mask = prev_seq >= 0 | ||
| 74 | + prev_dgamma = tl.load(dfactor_ptr + dfactor_offset + prev_seq * stride_dfactor_seq, mask=prev_mask, other=0.0) | ||
| 75 | + prev_dsgamma = prev_dgamma | ||
| 76 | + prev_strap = tl.load(trap_ptr + trap_offset + chunk_start * stride_trap_seq, mask=chunk_start < SEQLEN, other=0.0) | ||
| 77 | + prev_sdt = tl.load(dt_ptr + dt_offset + chunk_start * stride_dt_seq, mask=chunk_start < SEQLEN, other=0.0) | ||
| 78 | + prev_dsdt = tl.sigmoid(-prev_strap.to(tl.float32)) * prev_dsgamma | ||
| 79 | + prev_dstrap = -prev_sdt * prev_dsgamma | ||
| 80 | + | ||
| 81 | + offs_i = tl.arange(0, CHUNK_SIZE)[:, None] | ||
| 82 | + offs_j = tl.arange(0, CHUNK_SIZE)[None, :] | ||
| 83 | + shift_mask = offs_i == (offs_j + 1) | ||
| 84 | + dsdt_shift = tl.sum(tl.where(shift_mask, dsdt_block[None, :], 0.0), axis=1) | ||
| 85 | + dstrap_shift = tl.sum(tl.where(shift_mask, dstrap_block[None, :], 0.0), axis=1) | ||
| 86 | + | ||
| 87 | + offs = tl.arange(0, CHUNK_SIZE) | ||
| 88 | + dsdt_shift = tl.where(offs == 0, prev_dsdt, dsdt_shift) | ||
| 89 | + dstrap_shift = tl.where(offs == 0, prev_dstrap, dstrap_shift) | ||
| 90 | + | ||
| 91 | + ddt_out = dsdt_shift + dgamma_block * tl.sigmoid(trap_block.to(tl.float32)) | ||
| 92 | + dtrap_out = dstrap_shift + dgamma_block * dt_block | ||
| 93 | + dtrap_out *= tl.sigmoid(trap_block.to(tl.float32)) * tl.sigmoid(-trap_block.to(tl.float32)) | ||
| 94 | + | ||
| 95 | + ddt_ptrs = ddt_ptr + (pid_batch * stride_ddt_batch + pid_head * stride_ddt_head + offs_seq * stride_ddt_seq) | ||
| 96 | + dtrap_ptrs = dtrap_ptr + ( | ||
| 97 | + pid_batch * stride_dtrap_batch + pid_head * stride_dtrap_head + offs_seq * stride_dtrap_seq | ||
| 98 | + ) | ||
| 99 | + | ||
| 100 | + tl.store(ddt_ptrs, ddt_out, mask=offs_seq < SEQLEN) | ||
| 101 | + tl.store(dtrap_ptrs, dtrap_out, mask=offs_seq < SEQLEN) | ||
| @@ -0,0 +1,124 @@ | |||
| 1 | +api: pytorch | ||
| 2 | +version: v2.1 | ||
| 3 | +name: torch_bwd_dtrap_ddt | ||
| 4 | +triton_name: triton_bwd_dtrap_ddt.TritonBwdDtrapDdtFunctionApi | ||
| 5 | +api_type: torch_bwd_dtrap_ddt | ||
| 6 | +triton_api_type: triton_bwd_dtrap_ddt | ||
| 7 | +generate: generate_bwd_dtrap_ddt | ||
| 8 | +dtype_numbers: 20 | ||
| 9 | +standard: | ||
| 10 | + # bwd_dtrap_ddt is an elementwise backward op (sigmoid + mul + add) with no | ||
| 11 | + # matmul, so it is Vector-bound. The kernel does its intermediate math in | ||
| 12 | + # fp32 but the inputs/outputs stay in fp16/bf16, so the store-back rounding | ||
| 13 | + # dominates the error. With the narrow [-0.1, 0.1] input range the sigmoid | ||
| 14 | + # stays in its quasi-linear region and the measured NPU noise floor sits at | ||
| 15 | + # ~1e-4 (fp16) / ~1e-3 (bf16). 1e-3 sits ~10x above that floor and matches | ||
| 16 | + # the threshold used by the other mamba-derived backward siblings | ||
| 17 | + # (chunk_bwd_dqkwg, chunk_kda_bwd). | ||
| 18 | + acc: | ||
| 19 | + single_bm: | ||
| 20 | + type: high_performance | ||
| 21 | + fp16_error: 0.001 | ||
| 22 | + fp16_eb: 0.001 | ||
| 23 | + bf16_error: 0.001 | ||
| 24 | + bf16_eb: 0.001 | ||
| 25 | + fp32_error: 0.001 | ||
| 26 | + fp32_eb: 0.001 | ||
| 27 | + perf: not_key | ||
| 28 | +inputs: | ||
| 29 | + - name: trap | ||
| 30 | + type: tensor | ||
| 31 | + required: true | ||
| 32 | + dtypes: | ||
| 33 | + values: [ fp32, fp16, bf16 ] | ||
| 34 | + ranges: | ||
| 35 | + valid: | ||
| 36 | + values: [ [-0.1, 0.1] ] | ||
| 37 | + invalid: | ||
| 38 | + values: [ [-0.1, 0.1] ] | ||
| 39 | + shapes: | ||
| 40 | + dim_numbers: | ||
| 41 | + values: [ 3 ] | ||
| 42 | + dim_values: | ||
| 43 | + values: [ 1, 2, 4, 8, 16, 32 ] | ||
| 44 | + max_length: 131072 | ||
| 45 | + - name: dt | ||
| 46 | + type: tensor | ||
| 47 | + required: true | ||
| 48 | + dtypes: | ||
| 49 | + values: [ fp32, fp16, bf16 ] | ||
| 50 | + ranges: | ||
| 51 | + valid: | ||
| 52 | + values: [ [-0.1, 0.1] ] | ||
| 53 | + invalid: | ||
| 54 | + values: [ [-0.1, 0.1] ] | ||
| 55 | + shapes: | ||
| 56 | + dim_numbers: | ||
| 57 | + values: [ 3 ] | ||
| 58 | + dim_values: | ||
| 59 | + values: [ 1, 2, 4, 8, 16, 32 ] | ||
| 60 | + max_length: 131072 | ||
| 61 | + - name: dfactor | ||
| 62 | + type: tensor | ||
| 63 | + required: true | ||
| 64 | + dtypes: | ||
| 65 | + values: [ fp32, fp16, bf16 ] | ||
| 66 | + ranges: | ||
| 67 | + valid: | ||
| 68 | + values: [ [-0.1, 0.1] ] | ||
| 69 | + invalid: | ||
| 70 | + values: [ [-0.1, 0.1] ] | ||
| 71 | + shapes: | ||
| 72 | + dim_numbers: | ||
| 73 | + values: [ 3 ] | ||
| 74 | + dim_values: | ||
| 75 | + values: [ 1, 2, 4, 8, 16, 32 ] | ||
| 76 | + max_length: 131072 | ||
| 77 | + - name: dgamma_diag | ||
| 78 | + type: tensor | ||
| 79 | + required: true | ||
| 80 | + dtypes: | ||
| 81 | + values: [ fp32, fp16, bf16 ] | ||
| 82 | + ranges: | ||
| 83 | + valid: | ||
| 84 | + values: [ [-0.1, 0.1] ] | ||
| 85 | + invalid: | ||
| 86 | + values: [ [-0.1, 0.1] ] | ||
| 87 | + shapes: | ||
| 88 | + dim_numbers: | ||
| 89 | + values: [ 3 ] | ||
| 90 | + dim_values: | ||
| 91 | + values: [ 1, 2, 4, 8, 16, 32 ] | ||
| 92 | + max_length: 131072 | ||
| 93 | + - name: ddt | ||
| 94 | + type: tensor | ||
| 95 | + required: true | ||
| 96 | + dtypes: | ||
| 97 | + values: [ fp32, fp16, bf16 ] | ||
| 98 | + ranges: | ||
| 99 | + valid: | ||
| 100 | + values: [ [0, 0] ] | ||
| 101 | + invalid: | ||
| 102 | + values: [ [0, 0] ] | ||
| 103 | + shapes: | ||
| 104 | + dim_numbers: | ||
| 105 | + values: [ 3 ] | ||
| 106 | + dim_values: | ||
| 107 | + values: [ 1, 2, 4, 8, 16, 32 ] | ||
| 108 | + max_length: 131072 | ||
| 109 | + - name: dtrap | ||
| 110 | + type: tensor | ||
| 111 | + required: true | ||
| 112 | + dtypes: | ||
| 113 | + values: [ fp32, fp16, bf16 ] | ||
| 114 | + ranges: | ||
| 115 | + valid: | ||
| 116 | + values: [ [0, 0] ] | ||
| 117 | + invalid: | ||
| 118 | + values: [ [0, 0] ] | ||
| 119 | + shapes: | ||
| 120 | + dim_numbers: | ||
| 121 | + values: [ 3 ] | ||
| 122 | + dim_values: | ||
| 123 | + values: [ 1, 2, 4, 8, 16, 32 ] | ||
| 124 | + max_length: 131072 | ||
| @@ -0,0 +1,42 @@ | |||
| 1 | +# Copyright (c) 2023, Tri Dao, Albert Gu | ||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. | ||
| 3 | +# | ||
| 4 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | +# you may not use this file except in compliance with the License. | ||
| 6 | + | ||
| 7 | +from atk.case_generator.generator.base_generator import CaseGenerator | ||
| 8 | +from atk.case_generator.generator.generate_types import GENERATOR_REGISTRY | ||
| 9 | +from atk.configs.case_config import CaseConfig | ||
| 10 | + | ||
| 11 | + | ||
| 12 | +CASES = [ | ||
| 13 | + (1, 1, 8), | ||
| 14 | + (1, 1, 16), | ||
| 15 | + (1, 2, 32), | ||
| 16 | + (2, 4, 64), | ||
| 17 | + (1, 1, 128), | ||
| 18 | + (1, 8, 64), | ||
| 19 | + (2, 16, 256), | ||
| 20 | + (4, 32, 512), | ||
| 21 | +] | ||
| 22 | + | ||
| 23 | +CHUNK_SIZES = [8, 16, 32, 64] | ||
| 24 | + | ||
| 25 | + | ||
| 26 | + | ||
| 27 | +class BwdDtrapDdtGenerator(CaseGenerator): | ||
| 28 | + _case_index = 0 | ||
| 29 | + | ||
| 30 | + def after_case_config(self, case_config: CaseConfig) -> CaseConfig: | ||
| 31 | + idx = BwdDtrapDdtGenerator._case_index % len(CASES) | ||
| 32 | + B, H, S = CASES[idx] | ||
| 33 | + BwdDtrapDdtGenerator._case_index += 1 | ||
| 34 | + | ||
| 35 | + shape = [B, H, S] | ||
| 36 | + trap_dtype = case_config.inputs[0].dtype | ||
| 37 | + | ||
| 38 | + for i in range(len(case_config.inputs)): | ||
| 39 | + case_config.inputs[i].shape = shape | ||
| 40 | + case_config.inputs[i].dtype = trap_dtype | ||
| 41 | + | ||
| 42 | + return case_config | ||
| @@ -0,0 +1,135 @@ | |||
| 1 | +# Copyright (c) 2023, Tri Dao, Albert Gu | ||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. | ||
| 3 | +# | ||
| 4 | +# Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | +# you may not use this file except in compliance with the License. | ||
| 6 | + | ||
| 7 | +"""ATK adapter for bwd_dtrap_ddt. | ||
| 8 | + | ||
| 9 | +Inputs (positional, matches bwd_dtrap_ddt.yaml): | ||
| 10 | + args[0]: trap (B, H, S) | ||
| 11 | + args[1]: dt (B, H, S) | ||
| 12 | + args[2]: dfactor (B, H, S) | ||
| 13 | + args[3]: dgamma_diag (B, H, S) | ||
| 14 | + args[4]: ddt (B, H, S) output buffer, initialized to zeros | ||
| 15 | + args[5]: dtrap (B, H, S) output buffer, initialized to zeros | ||
| 16 | + | ||
| 17 | +Output: concatenated [ddt, dtrap] flattened. | ||
| 18 | +""" | ||
| 19 | + | ||
| 20 | +import torch | ||
| 21 | + | ||
| 22 | +from atk.configs.dataset_config import InputDataset | ||
| 23 | +from atk.tasks.api_execute import register | ||
| 24 | +from atk.tasks.api_execute.base_api import BaseApi | ||
| 25 | +from atk.tasks.api_execute.triton_base_api import TritonBaseApi | ||
| 26 | +from mindspeed_ops.api.triton.bwd_dtrap_ddt import bwd_dtrap_ddt | ||
| 27 | + | ||
| 28 | +CHUNK_SIZE = 64 | ||
| 29 | + | ||
| 30 | + | ||
| 31 | +def _get_inputs(input_data: InputDataset): | ||
| 32 | + keys = ("trap", "dt", "dfactor", "dgamma_diag", "ddt", "dtrap") | ||
| 33 | + if input_data.kwargs: | ||
| 34 | + return tuple(input_data.kwargs[key] for key in keys) | ||
| 35 | + return tuple(input_data.args[: len(keys)]) | ||
| 36 | + | ||
| 37 | + | ||
| 38 | +def _pack_outputs(outputs): | ||
| 39 | + return torch.cat([o.float().reshape(-1) for o in outputs]) | ||
| 40 | + | ||
| 41 | + | ||
| 42 | +def _process_chunk(trap_bh, dt_bh, dfactor_bh, dgamma_diag_bh, chunk_start, chunk_end, S, device): | ||
| 43 | + """Compute (ddt, dtrap) outputs for a single chunk.""" | ||
| 44 | + chunk_len = chunk_end - chunk_start | ||
| 45 | + | ||
| 46 | + trap_chunk = trap_bh[chunk_start:chunk_end] | ||
| 47 | + dt_chunk = dt_bh[chunk_start:chunk_end] | ||
| 48 | + dfactor_chunk = dfactor_bh[chunk_start:chunk_end] | ||
| 49 | + dgamma_diag_chunk = dgamma_diag_bh[chunk_start:chunk_end] | ||
| 50 | + | ||
| 51 | + # Build shifted strap/sdt: strap[i] = trap[chunk_start + i + 1] | ||
| 52 | + strap_chunk = torch.zeros(chunk_len, dtype=torch.float32, device=device) | ||
| 53 | + sdt_chunk = torch.zeros(chunk_len, dtype=torch.float32, device=device) | ||
| 54 | + valid_len = min(chunk_len, max(0, S - chunk_start - 1)) | ||
| 55 | + if valid_len > 0: | ||
| 56 | + strap_chunk[:valid_len] = trap_bh[chunk_start + 1 : chunk_start + 1 + valid_len] | ||
| 57 | + sdt_chunk[:valid_len] = dt_bh[chunk_start + 1 : chunk_start + 1 + valid_len] | ||
| 58 | + | ||
| 59 | + dgamma_chunk = dfactor_chunk + dgamma_diag_chunk | ||
| 60 | + dsdt_chunk = torch.sigmoid(-strap_chunk) * dfactor_chunk | ||
| 61 | + dstrap_chunk = -sdt_chunk * dfactor_chunk | ||
| 62 | + | ||
| 63 | + # Shift by 1 with chunk-boundary handling | ||
| 64 | + dsdt_shift = torch.zeros(chunk_len, dtype=torch.float32, device=device) | ||
| 65 | + dstrap_shift = torch.zeros(chunk_len, dtype=torch.float32, device=device) | ||
| 66 | + if chunk_start > 0: | ||
| 67 | + prev_dgamma = dfactor_bh[chunk_start - 1] | ||
| 68 | + dsdt_shift[0] = torch.sigmoid(-trap_bh[chunk_start]) * prev_dgamma | ||
| 69 | + dstrap_shift[0] = -dt_bh[chunk_start] * prev_dgamma | ||
| 70 | + if chunk_len > 1: | ||
| 71 | + dsdt_shift[1:] = dsdt_chunk[:-1] | ||
| 72 | + dstrap_shift[1:] = dstrap_chunk[:-1] | ||
| 73 | + | ||
| 74 | + sig_trap = torch.sigmoid(trap_chunk) | ||
| 75 | + sig_neg_trap = torch.sigmoid(-trap_chunk) | ||
| 76 | + ddt = dsdt_shift + dgamma_chunk * sig_trap | ||
| 77 | + dtrap = (dstrap_shift + dgamma_chunk * dt_chunk) * sig_trap * sig_neg_trap | ||
| 78 | + return ddt, dtrap | ||
| 79 | + | ||
| 80 | + | ||
| 81 | +def _bwd_dtrap_ddt_reference(trap, dt, dfactor, dgamma_diag, chunk_size): | ||
| 82 | + """Pure-fp32 CPU reference for bwd_dtrap_ddt.""" | ||
| 83 | + B, H, S = trap.shape | ||
| 84 | + | ||
| 85 | + ddt_out = torch.zeros(B, H, S, dtype=torch.float32, device=trap.device) | ||
| 86 | + dtrap_out = torch.zeros(B, H, S, dtype=torch.float32, device=trap.device) | ||
| 87 | + | ||
| 88 | + for b in range(B): | ||
| 89 | + for h in range(H): | ||
| 90 | + nchunks = (S + chunk_size - 1) // chunk_size | ||
| 91 | + for c in range(nchunks): | ||
| 92 | + chunk_start = c * chunk_size | ||
| 93 | + chunk_end = min(chunk_start + chunk_size, S) | ||
| 94 | + ddt_out[b, h, chunk_start:chunk_end], dtrap_out[b, h, chunk_start:chunk_end] = _process_chunk( | ||
| 95 | + trap[b, h], dt[b, h], dfactor[b, h], dgamma_diag[b, h], chunk_start, chunk_end, S, trap.device | ||
| 96 | + ) | ||
| 97 | + | ||
| 98 | + return ddt_out, dtrap_out | ||
| 99 | + | ||
| 100 | + | ||
| 101 | + | ||
| 102 | +class TorchBwdDtrapDdtFunctionApi(BaseApi): | ||
| 103 | + def __call__(self, input_data: InputDataset, with_output: bool = False): | ||
| 104 | + trap, dt, dfactor, dgamma_diag, _ddt, _dtrap = _get_inputs(input_data) | ||
| 105 | + trap_f = trap.float() | ||
| 106 | + dt_f = dt.float() | ||
| 107 | + dfactor_f = dfactor.float() | ||
| 108 | + dgamma_diag_f = dgamma_diag.float() | ||
| 109 | + | ||
| 110 | + ddt, dtrap = _bwd_dtrap_ddt_reference(trap_f, dt_f, dfactor_f, dgamma_diag_f, CHUNK_SIZE) | ||
| 111 | + return _pack_outputs((ddt, dtrap)) | ||
| 112 | + | ||
| 113 | + | ||
| 114 | + | ||
| 115 | +class TritonBwdDtrapDdtFunctionApi(TritonBaseApi): | ||
| 116 | + def __call__(self, input_data: InputDataset, with_output: bool = False): | ||
| 117 | + trap, dt, dfactor, dgamma_diag, ddt, dtrap = _get_inputs(input_data) | ||
| 118 | + | ||
| 119 | + B, H, S = trap.shape | ||
| 120 | + dtype = trap.dtype | ||
| 121 | + device = trap.device | ||
| 122 | + | ||
| 123 | + ddt_buf = torch.zeros(B, H, S, dtype=dtype, device=device) | ||
| 124 | + dtrap_buf = torch.zeros(B, H, S, dtype=dtype, device=device) | ||
| 125 | + | ||
| 126 | + bwd_dtrap_ddt( | ||
| 127 | + trap.contiguous(), | ||
| 128 | + dt.contiguous(), | ||
| 129 | + dfactor.contiguous(), | ||
| 130 | + dgamma_diag.contiguous(), | ||
| 131 | + ddt_buf, | ||
| 132 | + dtrap_buf, | ||
| 133 | + CHUNK_SIZE, | ||
| 134 | + ) | ||
| 135 | + return _pack_outputs((ddt_buf, dtrap_buf)) | ||
| @@ -0,0 +1,109 @@ | |||
| 1 | +# Copyright (c) 2026, HUAWEI CORPORATION. All rights reserved. | ||
| 2 | + | ||
| 3 | +import pytest | ||
| 4 | +import torch | ||
| 5 | + | ||
| 6 | +from mindspeed_ops.api.triton.bwd_dtrap_ddt import bwd_dtrap_ddt | ||
| 7 | +from mindspeed_ops.api.triton.utils import get_available_device | ||
| 8 | +from tests.utils import print_diff, assert_close | ||
| 9 | + | ||
| 10 | + | ||
| 11 | +def _dtype_thresholds(dtype): | ||
| 12 | + if dtype == torch.float32: | ||
| 13 | + return 1e-4, 1e-5 | ||
| 14 | + if dtype == torch.bfloat16: | ||
| 15 | + return 1e-2, 5e-3 | ||
| 16 | + if dtype == torch.float16: | ||
| 17 | + return 1e-3, 5e-4 | ||
| 18 | + raise ValueError(f"unsupported dtype {dtype}") | ||
| 19 | + | ||
| 20 | + | ||
| 21 | +TEST_SHAPES = [(1, 1, 8), (1, 1, 16), (1, 2, 32), (2, 4, 64), (1, 1, 128), (1, 8, 64)] | ||
| 22 | + | ||
| 23 | +CHUNK_SIZES = [8, 16, 32, 64] | ||
| 24 | + | ||
| 25 | +DTYPES = [torch.float32, torch.bfloat16, torch.float16] | ||
| 26 | + | ||
| 27 | + | ||
| 28 | +def _case_id(shape, chunk_size, dtype): | ||
| 29 | + B, H, S = shape | ||
| 30 | + dtype_name = {torch.float32: "fp32", torch.bfloat16: "bf16", torch.float16: "fp16"} | ||
| 31 | + return f"B{B}_H{H}_S{S}_C{chunk_size}_{dtype_name[dtype]}" | ||
| 32 | + | ||
| 33 | + | ||
| 34 | +TEST_CASES = [(shape, chunk_size, dtype) for shape in TEST_SHAPES for chunk_size in CHUNK_SIZES for dtype in DTYPES] | ||
| 35 | + | ||
| 36 | + | ||
| 37 | +def cpu_golden(trap, dt, dfactor, dgamma_diag, chunk_size): | ||
| 38 | + """Vectorized CPU reference for bwd_dtrap_ddt. | ||
| 39 | + | ||
| 40 | + The per-element computation at position p is: | ||
| 41 | + dsdt_shift[p] = sigmoid(-trap[p]) * dfactor[p-1] (0 for p=0) | ||
| 42 | + dstrap_shift[p] = -dt[p] * dfactor[p-1] (0 for p=0) | ||
| 43 | + ddt_out[p] = dsdt_shift[p] + (dfactor[p]+dgamma_diag[p]) * sigmoid(trap[p]) | ||
| 44 | + dtrap_out[p] = (dstrap_shift[p] + (dfactor[p]+dgamma_diag[p]) * dt[p]) * sigmoid(trap[p]) * sigmoid(-trap[p]) | ||
| 45 | + | ||
| 46 | + This is equivalent to the chunked loop version because the shift crosses chunk | ||
| 47 | + boundaries identically (dsdt at chunk_start uses dfactor[chunk_start-1]). | ||
| 48 | + """ | ||
| 49 | + dgamma = dfactor + dgamma_diag | ||
| 50 | + | ||
| 51 | + dfactor_prev = torch.zeros_like(dfactor) | ||
| 52 | + dfactor_prev[:, :, 1:] = dfactor[:, :, :-1] | ||
| 53 | + | ||
| 54 | + dsdt_shift = torch.sigmoid(-trap) * dfactor_prev | ||
| 55 | + dstrap_shift = -dt * dfactor_prev | ||
| 56 | + | ||
| 57 | + sig_trap = torch.sigmoid(trap) | ||
| 58 | + sig_neg_trap = torch.sigmoid(-trap) | ||
| 59 | + | ||
| 60 | + ddt_out = dsdt_shift + dgamma * sig_trap | ||
| 61 | + dtrap_out = (dstrap_shift + dgamma * dt) * sig_trap * sig_neg_trap | ||
| 62 | + | ||
| 63 | + return ddt_out, dtrap_out | ||
| 64 | + | ||
| 65 | + | ||
| 66 | +class TestBwdDtrapDdtOperator: | ||
| 67 | + def setup_method(self): | ||
| 68 | + torch.manual_seed(42) | ||
| 69 | + | ||
| 70 | + | ||
| 71 | + ("shape", "chunk_size", "dtype"), | ||
| 72 | + [pytest.param(*test, id=_case_id(*test)) for test in TEST_CASES], | ||
| 73 | + ) | ||
| 74 | + def test_bwd_dtrap_ddt_npu( | ||
| 75 | + self, | ||
| 76 | + shape, | ||
| 77 | + chunk_size: int, | ||
| 78 | + dtype: torch.dtype, | ||
| 79 | + ): | ||
| 80 | + B, H, S = shape | ||
| 81 | + | ||
| 82 | + trap = torch.randn(shape, dtype=dtype) | ||
| 83 | + dt = torch.randn(shape, dtype=dtype) | ||
| 84 | + dfactor = torch.randn(shape, dtype=dtype) | ||
| 85 | + dgamma_diag = torch.randn(shape, dtype=dtype) | ||
| 86 | + | ||
| 87 | + device = get_available_device() | ||
| 88 | + | ||
| 89 | + ddt_golden, dtrap_golden = cpu_golden( | ||
| 90 | + trap.float(), dt.float(), dfactor.float(), dgamma_diag.float(), chunk_size | ||
| 91 | + ) | ||
| 92 | + ddt_golden = ddt_golden.to(device).to(dtype) | ||
| 93 | + dtrap_golden = dtrap_golden.to(device).to(dtype) | ||
| 94 | + | ||
| 95 | + trap_npu = trap.to(device) | ||
| 96 | + dt_npu = dt.to(device) | ||
| 97 | + dfactor_npu = dfactor.to(device) | ||
| 98 | + dgamma_diag_npu = dgamma_diag.to(device) | ||
| 99 | + ddt_npu = torch.zeros(B, H, S, dtype=dtype, device=device) | ||
| 100 | + dtrap_npu = torch.zeros(B, H, S, dtype=dtype, device=device) | ||
| 101 | + | ||
| 102 | + bwd_dtrap_ddt(trap_npu, dt_npu, dfactor_npu, dgamma_diag_npu, ddt_npu, dtrap_npu, chunk_size) | ||
| 103 | + | ||
| 104 | + ratio, atol = _dtype_thresholds(dtype) | ||
| 105 | + print_diff("cpu_and_triton_ddt", ddt_golden, ddt_npu, atol) | ||
| 106 | + print_diff("cpu_and_triton_dtrap", dtrap_golden, dtrap_npu, atol) | ||
| 107 | + | ||
| 108 | + assert_close("cpu_and_triton_ddt", ddt_golden, ddt_npu, ratio, err_atol=atol) | ||
| 109 | + assert_close("cpu_and_triton_dtrap", dtrap_golden, dtrap_npu, ratio, err_atol=atol) | ||


需要补充GPU上pytorch小算子与GPU上开源triton算子的精度对齐截图(以确保pytorch小算子标杆的准确性)