Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ using namespace AscendC;
constexpr uint64_t BUFFER_NUM = 1;
constexpr uint32_t MAX_OUT_BUFFER_NUM = 2;
constexpr uint64_t MAX_MTP = 8;
constexpr uint64_t BF16_NUM_PER_BLOCK = 16;
constexpr uint64_t FP16_NUM_PER_BLOCK = 16;
constexpr uint64_t FP32_NUM_PER_BLOCK = 8;
constexpr uint32_t REPEAT_LENTH = 64; // 256Byte for float
constexpr uint32_t MAX_REPEAT_TIME = 255;
Expand Down Expand Up @@ -184,8 +184,8 @@ class RGDR {
stateOutBufferNum_ = (tilingData->stateOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM;
attnOutBufferNum_ = (tilingData->attnOutBufferNum == MAX_OUT_BUFFER_NUM) ? MAX_OUT_BUFFER_NUM : BUFFER_NUM;
restUbSize_ = tilingData->ubRestBytes;
alignK_ = Ceil(tilingData->dk, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK;
alignV_ = Ceil(tilingData->dv, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK;
alignK_ = Ceil(tilingData->dk, FP16_NUM_PER_BLOCK) * FP16_NUM_PER_BLOCK;
alignV_ = Ceil(tilingData->dv, FP16_NUM_PER_BLOCK) * FP16_NUM_PER_BLOCK;
load = 0;
usedblk = 0;
}
Expand Down Expand Up @@ -225,7 +225,7 @@ class RGDR {
uint32_t vSize = MAX_MTP * alignV_ * sizeof(float);
uint32_t kSize = MAX_MTP * alignK_ * sizeof(float);
uint32_t betaUbSize =
Ceil(MAX_MTP * NV_, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK * sizeof(float); // 8: 8 * 4 = 32B;
Ceil(MAX_MTP * NV_, FP16_NUM_PER_BLOCK) * FP16_NUM_PER_BLOCK * sizeof(float); // 8: 8 * 4 = 32B;
pipe_->InitBuffer(qInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType));
pipe_->InitBuffer(kInQueue_, BUFFER_NUM, MAX_MTP * alignK_ * sizeof(inType));
pipe_->InitBuffer(vInQueue_, BUFFER_NUM, MAX_MTP * alignV_ * sizeof(inType));
Expand Down Expand Up @@ -499,7 +499,7 @@ class RGDR {
__aicore__ inline void CopyInGamaBeta(int32_t seq0, int32_t seq1)
{
int32_t seqLen = seq1 - seq0;
uint64_t bBatchSize = Ceil(seqLen * NV_, BF16_NUM_PER_BLOCK) * BF16_NUM_PER_BLOCK;
uint64_t bBatchSize = Ceil(seqLen * NV_, FP16_NUM_PER_BLOCK) * FP16_NUM_PER_BLOCK;
LocalTensor<inType> betaLocal = betaInQueue_.AllocTensor<inType>();
DataCopyParams betaInParams{1, static_cast<uint16_t>(seqLen * NV_ * sizeof(inType)), 0, 0};
DataCopyCustom(betaLocal, betaGm_[seq0 * NV_], betaInParams);
Expand Down
10 changes: 5 additions & 5 deletions csrc/moe/causal_conv1d_v310/causal_conv1d_310_torch_adpt.h
Original file line number Diff line number Diff line change
Expand Up @@ -22,10 +22,10 @@ at::Tensor npu_causal_conv1d_310(
const at::Tensor& weight,
const c10::optional<at::Tensor>& bias,
const at::Tensor& conv_states,
at::IntArrayRef query_start_loc,
at::IntArrayRef cache_indices,
at::IntArrayRef initial_state_mode,
at::IntArrayRef num_accepted_tokens,
const c10::optional<at::Tensor>& query_start_loc,
const c10::optional<at::Tensor>& cache_indices,
const c10::optional<at::Tensor>& initial_state_mode,
const c10::optional<at::Tensor>& num_accepted_tokens,
int64_t activation_mode,
int64_t pad_slot_id,
int64_t run_mode)
Expand All @@ -50,4 +50,4 @@ at::Tensor npu_causal_conv1d_310(
}

}
#endif
#endif
Original file line number Diff line number Diff line change
Expand Up @@ -89,4 +89,4 @@ class CausalConv1dV310 : public OpDef {
};
OP_ADD(CausalConv1dV310);

} // namespace ops
} // namespace ops
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@ static inline DimTileChoice ChooseDimTileSize(gert::TilingContext *context, int6
const int64_t candidates[] = {4096, 2048, 1024, 512, 384, 192};

auto ChooseOnce = [&](bool requireExactDiv) -> DimTileChoice {
const bool preferMoreBlocks = (batch > static_cast<int64_t>(coreNum));
DimTileChoice bestOver;
int64_t bestOverGap = std::numeric_limits<int64_t>::max();
DimTileChoice bestUnder;
Expand All @@ -79,7 +80,11 @@ static inline DimTileChoice ChooseDimTileSize(gert::TilingContext *context, int6
if (requireExactDiv && (dim % dimTileSize != 0)) {
continue;
}
const int64_t blocksPerSeq = requireExactDiv ? (dim / dimTileSize) : CeilDivInt64(dim, dimTileSize);
const int64_t testBlocksPerSeq = requireExactDiv ? (dim / dimTileSize) : CeilDivInt64(dim, dimTileSize);
if (preferMoreBlocks && testBlocksPerSeq <= 1) {
continue;
}
const int64_t blocksPerSeq = testBlocksPerSeq;
const int64_t gridSize = batch * blocksPerSeq;
if (gridSize <= 0) {
continue;
Expand All @@ -89,7 +94,6 @@ static inline DimTileChoice ChooseDimTileSize(gert::TilingContext *context, int6
if (gridSize >= static_cast<int64_t>(coreNum)) {
const int64_t gap = gridSize - static_cast<int64_t>(coreNum);
if (gap < bestOverGap) {
// bestOver = {dimTileSize, blocksPerSeq, gridSize};
bestOver.dimTileSize = dimTileSize;
bestOver.blocksPerSeq = blocksPerSeq;
bestOver.gridSize = gridSize;
Expand Down Expand Up @@ -585,4 +589,4 @@ static ge::graphStatus TilingParseForCausalConv1d(gert::TilingParseContext *cont
IMPL_OP_OPTILING(CausalConv1dV310)
.Tiling(CausalConv1dTilingFunc)
.TilingParse<CausalConv1dCompileInfo>(TilingParseForCausalConv1d);
} // namespace optiling
} // namespace optiling
9 changes: 7 additions & 2 deletions csrc/moe/causal_conv1d_v310/op_kernel/causal_conv1d_v310.h
Original file line number Diff line number Diff line change
Expand Up @@ -338,6 +338,11 @@ __aicore__ inline void CausalConv1dV310<T>::WriteBackState(int32_t cacheIdx, int
if (len <= 0) {
return;
}
const int32_t expectedDimTileSize = (c0 + dimTileSize <= dim) ? dimTileSize : (dim - c0);
if (c0 < 0 || c0 >= dim || dimTileSize <= 0 || dimTileSize > expectedDimTileSize) {
// Invalid c0 or dimTileSize would cause wrong state writeback
return;
}

const int32_t lastT = len - 1;
LocalTensor<T> ring = inBuf.Get<T>();
Expand Down Expand Up @@ -495,7 +500,7 @@ __aicore__ inline void CausalConv1dV310<T>::Process()
cacheIdx = static_cast<int32_t>(cacheIdx64);
}

const bool hasInit = (tilingData_->hasInitialStateMode != 0) ? (initialStateModeGm.GetValue(seq) != 0) : false;
const bool hasInit = (tilingData_->hasInitialStateMode != 0) ? (initialStateModeGm.GetValue(seq) != 0) : true;
int32_t stateTokenOffset = 0;
if (isSpecDecodingGlobal) {
int32_t accepted = static_cast<int32_t>(numAcceptedTokensGm.GetValue(seq));
Expand Down Expand Up @@ -539,4 +544,4 @@ __aicore__ inline void CausalConv1dV310<T>::Process()
}

} // namespace NsCausalConv1d
#endif // CAUSAL_CONV1D_V310_H
#endif // CAUSAL_CONV1D_V310_H
8 changes: 4 additions & 4 deletions csrc/torch_binding.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1886,10 +1886,10 @@ TORCH_LIBRARY_EXPAND(CONCAT(_C, _ascend), ops)
" Tensor weight, "
" Tensor? bias, "
" Tensor conv_states, "
" int[] query_start_loc, "
" int[] cache_indices, "
" int[] initial_state_mode, "
" int[] num_accepted_tokens, "
" Tensor? query_start_loc, "
" Tensor? cache_indices, "
" Tensor? initial_state_mode, "
" Tensor? num_accepted_tokens, "
" int activation_mode, "
" int pad_slot_id, "
" int run_mode) -> (Tensor output)");
Comment on lines 1886 to 1895

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

Suggested PR Title:

[310p][Ops][Feature] use recurrent gdn custom op with aclgraph

Suggested PR Summary:

### What this PR does / why we need it?
This PR introduces a custom recurrent GDN operator for Ascend 310P to improve performance and support `aclgraph` capture. It replaces the existing PyTorch-based recurrent rule with a fused kernel (`npu_recurrent_gated_delta_rule_310`) and updates the `causal_conv1d_310` operator to use tensors for metadata instead of array refs, facilitating graph capture.

### Does this PR introduce _any_ user-facing change?
No. This is an internal performance optimization for Ascend 310P.

### How was this patch tested?
Tested with existing E2E nightly tests for `causal_conv1d_310` and GDN attention.

The current PR title and summary do not follow the required format specified in the repository style guide. I have provided a suggested title and summary above.

Additionally, the schema for npu_causal_conv1d_310 is missing the in-place marker for the conv_states tensor. Since this operator updates the convolution states in-place (as verified by the tests), the schema should use Tensor! conv_states to ensure correct behavior with the PyTorch dispatcher and graph capture mechanisms.

        "                         Tensor weight, "
        "                         Tensor? bias, "
        "                         Tensor! conv_states, "
        "                         Tensor? query_start_loc, "
        "                         Tensor? cache_indices, "
        "                         Tensor? initial_state_mode, "
        "                         Tensor? num_accepted_tokens, "
        "                         int activation_mode, "
        "                         int pad_slot_id, "
        "                         int run_mode) -> (Tensor output)");
References
  1. The PR title and summary must follow the specific format: [Branch][Module][Action] Title, and include sections for 'What this PR does', 'User-facing change', and 'How was this patch tested'. (link)

Expand Down
8 changes: 4 additions & 4 deletions csrc/torch_binding_meta.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -601,10 +601,10 @@ at::Tensor npu_causal_conv1d_310_meta(
const at::Tensor& weight,
const c10::optional<at::Tensor>& bias,
const at::Tensor& conv_states,
at::IntArrayRef query_start_loc,
at::IntArrayRef cache_indices,
at::IntArrayRef initial_state_mode,
at::IntArrayRef num_accepted_tokens,
const c10::optional<at::Tensor>& query_start_loc,
const c10::optional<at::Tensor>& cache_indices,
const c10::optional<at::Tensor>& initial_state_mode,
const c10::optional<at::Tensor>& num_accepted_tokens,
int64_t activation_mode,
int64_t pad_slot_id,
int64_t run_mode)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,13 +16,6 @@ def validate_cmp(y_cal, y_ref, device="npu"):
torch.testing.assert_close(y_ref, y_cal, rtol=3e-03, atol=1e-02, equal_nan=True)


def to_int64_tuple(t):
t = t.to(torch.int64)
if t.dim() == 0:
return (t.item(),)
return tuple(t.tolist())


@pytest.mark.parametrize("has_initial_state", [False, True])
@pytest.mark.parametrize("silu_activation", [True])
@pytest.mark.parametrize("has_bias", [True])
Expand Down Expand Up @@ -85,10 +78,10 @@ def test_ascend_causal_conv1d_310_fn(
weight_origin,
bias=bias,
conv_states=conv_states_origin,
query_start_loc=to_int64_tuple(query_start_loc),
cache_indices=to_int64_tuple(cache_indices),
initial_state_mode=to_int64_tuple(has_initial_state_tensor),
num_accepted_tokens=[],
query_start_loc=query_start_loc.to(torch.int64),
cache_indices=cache_indices.to(torch.int64),
initial_state_mode=has_initial_state_tensor.to(torch.int64),
num_accepted_tokens=None,
activation_mode=activation_mode,
pad_slot_id=PAD_SLOT_ID,
run_mode=0,
Expand Down Expand Up @@ -130,17 +123,16 @@ def test_causal_conv1d_310_update(batch_size, dim, width, seqlen, has_bias, silu
activation = None if not silu_activation else "silu"

activation_mode = 1 if activation else 0
has_initial_state_tensor = torch.tensor([True] * batch_size, device=device, dtype=torch.bool)
conv_states_origin = conv_states.transpose(-1, -2)
out = torch.ops._C_ascend.npu_causal_conv1d_310(
x.transpose(-1, -2),
weight.transpose(-1, -2),
bias=bias,
conv_states=conv_states_origin,
query_start_loc=[],
cache_indices=to_int64_tuple(conv_state_indices),
initial_state_mode=to_int64_tuple(has_initial_state_tensor),
num_accepted_tokens=[],
query_start_loc=None,
cache_indices=conv_state_indices.to(torch.int64),
initial_state_mode=None,
num_accepted_tokens=None,
activation_mode=activation_mode,
pad_slot_id=PAD_SLOT_ID,
run_mode=1,
Expand Down
Loading
Loading