From 50cb7a62205296bb718cb5006cea3d945ec3d2e0 Mon Sep 17 00:00:00 2001 From: ChaoZheng109 Date: Fri, 28 Aug 2026 18:51:13 -0700 Subject: [PATCH] Add: an AUTO-scope twin of the qwen3_14b_decode scene test qwen3_14b_decode runs under ScopeMode::MANUAL, which returns from compute_task_fanin immediately and bypasses TensorMap entirely. It exercises no dependency derivation and cannot measure a change to it, so no large-model workload covers the AUTO path. Add qwen3_14b_decode_auto: the same network, fixture and golden under ScopeMode::AUTO, where every WAIT edge is derived from tensor overlap. The two directories are an A/B pair on one workload. The sibling is unchanged, and only the 13 incores that had to differ live in the new directory -- the CALLABLE names the other 27, including the vendored CANN attention tree, as ../qwen3_14b_decode/kernels/..., so the harvested codegen keeps one home. Switching scope mode alone costs 6.3x, and the cause is not the derivation but two things the codegen never had to express. The first is declaration width. A task that declares a whole buffer while writing one column band makes TensorMap order it against every other band; down_acc_all's 17 k-splits each write their own 1024-wide band and were serialized by the parent declaration. Narrowing each argument to the band its kernel indexes moves the argument's base, so the kernels that lost a column term did so to index relative to the view. out_proj_0 keeps its parent declaration because it resolves its band from block_idx, and there is no constant to drop. The second is commutative accumulation. AtomicAdd into a shared band does not care about order, but INOUT cannot say so and TensorMap must serialize the writers. down_proj, out_proj, gate_proj_4 and up_proj_4 now write private partials that a partials_reduce_* task sums, turning an N-deep chain into one parallel round plus one task. out_proj's reducer accumulates rather than stores, because out_proj_0's SPMD blocks write the same bands from a block_idx-derived offset; its split count is per-band, since the direct loop stops at N_OUT_DIRECT and summing a fixed five would fold in slabs no task wrote. residual_rms_cast and its four variants needed only the first: they write disjoint bands with plain stores, so narrowing their two outputs and indexing them relative to the band removes the chain outright. Critical path over the WAIT subgraph, one decode step on a2a3: MANUAL sibling 443 AUTO, parent declarations 7,963 AUTO, banded declarations 1,563 AUTO, plus private partials 684 Edge count rises from 58,404 to 87,938 across those rows: narrowing a declaration replaces one long-range edge with several short-range ones, so edge count is not a proxy for parallelism. Of the remaining 241 steps, half are gate_proj and gate_proj_0..3 -- five separate SPMD tasks accumulating into one column range, curable the same way at the cost of ten near-duplicate kernels. The other half is the reducers themselves: MANUAL declares by hand that its 85 down_proj tasks are independent and pays nothing, where AUTO's floor is two edges. Closing that needs an argument direction marking a write commutative. Both cases pass on a2a3 at RTOL=5e-2 / ATOL=1e-1 across the output and all 40 layers' KV caches. --- .../qwen3_14b_decode_auto/README.md | 107 + .../kernels/aic/down_proj.cpp | 620 +++++ .../kernels/aic/gate_proj_4.cpp | 621 +++++ .../kernels/aic/out_proj.cpp | 576 +++++ .../kernels/aic/up_proj_4.cpp | 621 +++++ .../kernels/aiv/partials_reduce.h | 70 + .../kernels/aiv/partials_reduce_hidden.cpp | 69 + .../aiv/partials_reduce_hidden_add.cpp | 70 + .../kernels/aiv/partials_reduce_inter.cpp | 69 + .../kernels/aiv/residual_rms_cast.cpp | 369 +++ .../kernels/aiv/residual_rms_cast_0.cpp | 369 +++ .../kernels/aiv/residual_rms_cast_1.cpp | 369 +++ .../kernels/aiv/residual_rms_cast_2.cpp | 369 +++ .../kernels/aiv/residual_rms_cast_3.cpp | 369 +++ .../kernels/aiv/silu.cpp | 422 +++ .../orchestration/decode_fwd_layers.cpp | 2265 +++++++++++++++++ .../test_qwen3_14b_decode_auto.py | 470 ++++ 17 files changed, 7825 insertions(+) create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/README.md create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/down_proj.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/gate_proj_4.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/out_proj.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/up_proj_4.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce.h create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden_add.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_inter.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_0.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_1.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_2.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_3.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/silu.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/orchestration/decode_fwd_layers.cpp create mode 100644 examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/test_qwen3_14b_decode_auto.py diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/README.md b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/README.md new file mode 100644 index 0000000000..5084491b7d --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/README.md @@ -0,0 +1,107 @@ +# `qwen3_14b_decode_auto/` — the TensorMap-derived-dependency twin of `qwen3_14b_decode/` + +Same network, same numerics, same fixture as +[`qwen3_14b_decode/`](../qwen3_14b_decode/README.md) — read that README for the +model, the parameter regime, the attention extern, provenance and cost. This +directory exists only so the two dependency-derivation modes can be measured +against each other on one workload. + +The sibling declares every dependency by hand inside +`SIMPLER_SCOPE(ScopeMode::MANUAL)`, which returns from `compute_task_fanin` +immediately and switches TensorMap off entirely. It therefore exercises no +dependency derivation at all. This one runs under `ScopeMode::AUTO`, so every +WAIT edge is derived from tensor overlap. + +## What is actually here + +Everything identical to the sibling is **shared, not copied**: the `CALLABLE` +names 27 of its 41 incores as `../qwen3_14b_decode/kernels/…`, so a refresh of +the harvested codegen lands in one place. Only the files that had to change are +local: + +| file | why it differs | +| ---- | -------------- | +| `kernels/orchestration/decode_fwd_layers.cpp` | `ScopeMode::AUTO`; banded views; private partials; reducer submits | +| `kernels/aic/{down,out}_proj.cpp` | stores its own partial instead of accumulating | +| `kernels/aic/{gate,up}_proj_4.cpp` | stores its own partial instead of accumulating | +| `kernels/aiv/silu.cpp` | column term dropped from the base pointer | +| `kernels/aiv/residual_rms_cast{,_0,_1,_2,_3}.cpp` | its two stores index relative to the band | +| `kernels/aiv/partials_reduce*.{h,cpp}` | new; sums private split-K partials | + +## The two things AUTO needs that MANUAL does not + +### 1. Declare each view at the width the task actually touches + +A task that declares a whole buffer while writing one column band makes +TensorMap order it against every other band. `down_acc_all` is `[16, 17408]`; +17 k-splits each write their own 1024-wide band but the harvested codegen +declared the parent, so all 17 were serialized. The `.slice()` calls here +narrow each argument to the band its kernel indexes. + +**A narrowed view moves the argument's base**, so this only works when the +kernel indexes relative to it. Every kernel listed above that lost a column term +did so for this reason. Two places therefore keep the parent declaration: +`out_proj_0` resolves its band as `(idx / 5) * 512` measured from the buffer +base, and there is no constant to drop because `idx` is `block_idx`-derived. + +### 2. Give commutative accumulation somewhere to land + +`AtomicAdd` into a shared band is commutative, but `INOUT` cannot say so, so +TensorMap must serialize the writers. Four accumulators — `down_proj` (17 +splits), `out_proj`, `gate_proj_4` and `up_proj_4` (5 each) — now write private +partials that a `partials_reduce_*` task sums, which turns an N-deep chain into +one parallel round plus one task. + +`out_proj`'s reducer accumulates rather than stores: `out_proj_0`'s SPMD blocks +write the same bands from a `block_idx`-derived offset, so the sum has to join +them. Its split count is per-band, because the direct loop stops at +`N_OUT_DIRECT` and the last band is short — summing a fixed 5 would fold in +slabs no task wrote. + +## Measured + +One decode step on a2a3, same die; WAIT-edge graph from `--enable-dep-gen`, +critical path = longest path over the WAIT subgraph. + +| variant | critical path | per layer | edges | +| ------- | ------------: | --------: | ----: | +| MANUAL sibling | 443 | 11.1 | 23,605 | +| AUTO, parent declarations everywhere | 7,963 | 199.1 | 58,404 | +| AUTO, banded declarations | 1,563 | 39.1 | 66,604 | +| **AUTO, + private split-K partials (this)** | **684** | **17.1** | 87,938 | + +Edge count rises as the critical path falls: narrowing a declaration replaces +one long-range edge with several short-range ones. Edge count is not a proxy for +parallelism — the critical path is. + +## What the remaining 241 steps are + +Both modes walk the same per-layer skeleton; AUTO takes 17 edges where MANUAL +takes 11, and the 6 extra split evenly: + +- **120 steps — `gate_proj`, `gate_proj_0..3`.** Five separate SPMD tasks, one + per k-split, all accumulating into columns `[0, 6144)`. The six blocks *inside* + each task run in parallel; the five tasks are what serialize. Curable by the + same private-partial treatment, at the cost of ten near-duplicate kernels whose + only difference from the sibling's would be `AtomicAdd` → `AtomicNone`. +- **120 steps — the reducers themselves.** MANUAL gets atomic accumulation for + free: it declares by hand that the 85 `down_proj` tasks are mutually + independent, so accumulation costs zero critical-path steps. AUTO's best is two + edges — one parallel round, then the reducer. Closing this needs an argument + direction that marks a write commutative, so TensorMap can leave the writers + unordered; no amount of slicing reaches it. + +## Running + +```bash +pytest examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto \ + --platform a2a3 --device ${DEVICE} --manual include + +# the A/B: add --enable-dep-gen and run the sibling the same way +pytest examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto \ + --platform a2a3 --device ${DEVICE} --manual include --enable-dep-gen +``` + +Like the sibling it runs in the daily full scene-test sweep, not per-PR CI, and +passes at `RTOL=5e-2 / ATOL=1e-1` against the same golden — output and all 40 +layers' KV caches. diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/down_proj.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/down_proj.cpp new file mode 100644 index 0000000000..13f24f2123 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/down_proj.cpp @@ -0,0 +1,620 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: down_proj + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +static __aicore__ void +down_proj(__gm__ bfloat16_t *v1, __gm__ bfloat16_t *v2, __gm__ float *v3, int64_t v4, int64_t v5, int64_t v6) { + const int64_t v7 = 960; + const int64_t v8 = 2; + const int64_t v9 = 15; + const int64_t v10 = 32; + const int64_t v11 = 64; + const int64_t v12 = 5120; + const int64_t v13 = 1; + const int64_t v14 = 17408; + const int64_t v15 = 16; + const int64_t v16 = 512; + const int64_t v17 = 2048; + const int64_t v18 = 32768; + const int64_t v19 = 1536; + const int64_t v20 = 1024; + const int64_t v21 = 0; + const int64_t v22 = 135168; + const int64_t v23 = 4096; + using T = float; + +#if defined(__DAV_CUBE__) + size_t v24 = (size_t)v11; + size_t v25 = (size_t)v21; + size_t v26 = (size_t)v10; + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v27 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v28 = (uint64_t)v23; + TASSIGN(v27, v28); + pto::Shape<1, 1, 1, 16, 64> v29 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v30 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v31 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14), v29, v30 + ); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + TLOAD(v27, v31); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v32 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v33 = (uint64_t)v22; + TASSIGN(v32, v33); + int64_t v34 = (int64_t)((uint64_t)v5 + (uint64_t)v4); + pto::Shape<1, 1, 1, 64, 1024> v35 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v36 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v37 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND>( + v2 + (v21 + v34 * v12 + v6 * v13), v35, v36 + ); + TLOAD(v32, v37); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + for (size_t v38 = v25; v38 < v24; v38 += v26) { + int64_t v39 = (int64_t)v38; + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v40 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v41 = (uint64_t)v20; + TASSIGN(v40, v41); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + pipe_barrier(PIPE_MTE1); + TEXTRACT(v40, v27, v21, v38); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v42 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v43 = (uint64_t)v21; + TASSIGN(v42, v43); + TEXTRACT(v42, v32, v38, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v44 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v45 = (uint64_t)v19; + TASSIGN(v44, v45); + int64_t v46 = (int64_t)((uint64_t)v39 + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v44, v27, v21, v46); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v47 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v48 = (uint64_t)v18; + TASSIGN(v47, v48); + TEXTRACT(v47, v32, v46, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + if (v39 == v21) { + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v49 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v50 = (uint64_t)v21; + TASSIGN(v49, v50); + pipe_barrier(PIPE_M); + TMATMUL(v49, v40, v42); + } else { + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v52 = (uint64_t)v21; + TASSIGN(v51, v52); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v51, v51, v40, v42); + }; + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v53 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v54 = (uint64_t)v21; + TASSIGN(v53, v54); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + TMATMUL_ACC(v53, v53, v44, v47); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + for (size_t v55 = (size_t)v13; v55 < ((size_t)v9); v55 += (size_t)v8) { + int64_t v56 = (int64_t)((uint64_t)((int64_t)v55) * (uint64_t)v11); + int64_t v57 = (int64_t)((uint64_t)v56 + (uint64_t)v11); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v59 = (uint64_t)v21; + TASSIGN(v58, v59); + pto::Shape<1, 1, 1, 16, 64> v60 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v61 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v62 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, + pto::Layout::ND>(v1 + (v21 + v21 * v14 + (int64_t)(uint64_t)v56 * v13), v60, v61); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + TLOAD(v58, v62); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v63 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v64 = (uint64_t)v22; + TASSIGN(v63, v64); + pto::Shape<1, 1, 1, 64, 1024> v65 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v66 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v56) * v12 + v6 * v13), v65, v66); + TLOAD(v63, v67); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v69 = (uint64_t)v17; + TASSIGN(v68, v69); + pto::Shape<1, 1, 1, 16, 64> v70 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v71 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v72 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, + pto::Layout::ND>(v1 + (v21 + v21 * v14 + (int64_t)(uint64_t)v57 * v13), v70, v71); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + TLOAD(v68, v72); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v73 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v74 = (uint64_t)v23; + TASSIGN(v73, v74); + pto::Shape<1, 1, 1, 64, 1024> v75 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v76 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v77 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v57) * v12 + v6 * v13), v75, v76); + TLOAD(v73, v77); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + for (size_t v78 = v25; v78 < v24; v78 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v79 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v80 = (uint64_t)v20; + TASSIGN(v79, v80); + pipe_barrier(PIPE_MTE1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + TEXTRACT(v79, v58, v21, v78); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v81 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v82 = (uint64_t)v21; + TASSIGN(v81, v82); + TEXTRACT(v81, v63, v78, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v83 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v84 = (uint64_t)v19; + TASSIGN(v83, v84); + int64_t v85 = (int64_t)((uint64_t)((int64_t)v78) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + TEXTRACT(v83, v58, v21, v85); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v86 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v87 = (uint64_t)v18; + TASSIGN(v86, v87); + TEXTRACT(v86, v63, v85, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v88 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v89 = (uint64_t)v21; + TASSIGN(v88, v89); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v88, v88, v79, v81); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v90 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v91 = (uint64_t)v21; + TASSIGN(v90, v91); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + TMATMUL_ACC(v90, v90, v83, v86); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + }; + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID6); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID6); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + for (size_t v92 = v25; v92 < v24; v92 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v93 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v94 = (uint64_t)v21; + TASSIGN(v93, v94); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + TEXTRACT(v93, v68, v21, v92); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v95 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v96 = (uint64_t)v21; + TASSIGN(v95, v96); + TEXTRACT(v95, v73, v92, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v97 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v98 = (uint64_t)v16; + TASSIGN(v97, v98); + int64_t v99 = (int64_t)((uint64_t)((int64_t)v92) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + TEXTRACT(v97, v68, v21, v99); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v100 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v101 = (uint64_t)v18; + TASSIGN(v100, v101); + TEXTRACT(v100, v73, v99, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v102 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v103 = (uint64_t)v21; + TASSIGN(v102, v103); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v102, v102, v93, v95); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v104 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v105 = (uint64_t)v21; + TASSIGN(v104, v105); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + TMATMUL_ACC(v104, v104, v97, v100); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + }; + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v106 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v107 = (uint64_t)v23; + TASSIGN(v106, v107); + pto::Shape<1, 1, 1, 16, 64> v108 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v109 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v110 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + (int64_t)(uint64_t)v7 * v13), v108, v109 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + TLOAD(v106, v110); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v111 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v112 = (uint64_t)v22; + TASSIGN(v111, v112); + pto::Shape<1, 1, 1, 64, 1024> v113 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v114 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v115 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND>( + v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v7) * v12 + v6 * v13), v113, v114 + ); + TLOAD(v111, v115); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + for (size_t v116 = v25; v116 < v24; v116 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v117 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v118 = (uint64_t)v20; + TASSIGN(v117, v118); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + TEXTRACT(v117, v106, v21, v116); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v119 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v120 = (uint64_t)v21; + TASSIGN(v119, v120); + TEXTRACT(v119, v111, v116, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v121 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v122 = (uint64_t)v19; + TASSIGN(v121, v122); + int64_t v123 = (int64_t)((uint64_t)((int64_t)v116) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v121, v106, v21, v123); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v124 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v125 = (uint64_t)v18; + TASSIGN(v124, v125); + TEXTRACT(v124, v111, v123, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v126 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v127 = (uint64_t)v21; + TASSIGN(v126, v127); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v126, v126, v117, v119); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v128 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v129 = (uint64_t)v21; + TASSIGN(v128, v129); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + TMATMUL_ACC(v128, v128, v121, v124); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v130 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v131 = (uint64_t)v21; + TASSIGN(v130, v131); + pto::Shape<1, 1, 1, 16, 1024> v132 = pto::Shape<1, 1, 1, 16, 1024>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v133 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v134 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 1024>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v21 + v21 * v12), v132, v133 + ); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + TSTORE< + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>, + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>, + AtomicType::AtomicNone>(v134, v130); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); +#endif // __DAV_CUBE__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_tile_inline149__rv_v2 + __gm__ Tensor *mlp_tile_inline149__rv_v2_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_tile_inline149__rv_v2 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_tile_inline149__rv_v2_tensor->buffer.addr) + + mlp_tile_inline149__rv_v2_tensor->start_offset; + + // Unpack tensor: w_down__ssa_v0 + __gm__ Tensor *w_down__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ bfloat16_t *w_down__ssa_v0 = + reinterpret_cast<__gm__ bfloat16_t *>(w_down__ssa_v0_tensor->buffer.addr) + w_down__ssa_v0_tensor->start_offset; + + // Unpack tensor: down_acc_all_inline168__iter_v6 + __gm__ Tensor *down_acc_all_inline168__iter_v6_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *down_acc_all_inline168__iter_v6 = + reinterpret_cast<__gm__ float *>(down_acc_all_inline168__iter_v6_tensor->buffer.addr) + + down_acc_all_inline168__iter_v6_tensor->start_offset; + + // Unpack scalar: k0_inline113__ssa_v8 + union { + uint64_t u64; + int64_t val; + } k0_inline113__ssa_v8_conv; + k0_inline113__ssa_v8_conv.u64 = args[3]; + int64_t k0_inline113__ssa_v8 = k0_inline113__ssa_v8_conv.val; + + // Unpack scalar: layer_inter_base_inline107__ssa_v0 + union { + uint64_t u64; + int64_t val; + } layer_inter_base_inline107__ssa_v0_conv; + layer_inter_base_inline107__ssa_v0_conv.u64 = args[4]; + int64_t layer_inter_base_inline107__ssa_v0 = layer_inter_base_inline107__ssa_v0_conv.val; + + // Unpack scalar: n0_inline122__ssa_v8 + union { + uint64_t u64; + int64_t val; + } n0_inline122__ssa_v8_conv; + n0_inline122__ssa_v8_conv.u64 = args[5]; + int64_t n0_inline122__ssa_v8 = n0_inline122__ssa_v8_conv.val; + + // Forward to ptoas-generated function + down_proj( + mlp_tile_inline149__rv_v2, w_down__ssa_v0, down_acc_all_inline168__iter_v6, k0_inline113__ssa_v8, + layer_inter_base_inline107__ssa_v0, n0_inline122__ssa_v8 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/gate_proj_4.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/gate_proj_4.cpp new file mode 100644 index 0000000000..82079f9495 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/gate_proj_4.cpp @@ -0,0 +1,621 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: gate_proj_4 + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +static __aicore__ void +gate_proj_4(__gm__ bfloat16_t *v1, __gm__ bfloat16_t *v2, __gm__ float *v3, int64_t v4, int64_t v5, int64_t v6) { + const int64_t v7 = 960; + const int64_t v8 = 2; + const int64_t v9 = 15; + const int64_t v10 = 32; + const int64_t v11 = 64; + const int64_t v12 = 17408; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 512; + const int64_t v17 = 2048; + const int64_t v18 = 32768; + const int64_t v19 = 1536; + const int64_t v20 = 1024; + const int64_t v21 = 0; + const int64_t v22 = 135168; + const int64_t v23 = 4096; + using T = float; + +#if defined(__DAV_CUBE__) + size_t v24 = (size_t)v11; + size_t v25 = (size_t)v21; + size_t v26 = (size_t)v10; + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v27 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v28 = (uint64_t)v23; + TASSIGN(v27, v28); + pto::Shape<1, 1, 1, 16, 64> v29 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v30 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v31 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + v4 * v13), v29, v30 + ); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + TLOAD(v27, v31); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v32 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v33 = (uint64_t)v22; + TASSIGN(v32, v33); + int64_t v34 = (int64_t)((uint64_t)v5 + (uint64_t)v4); + pto::Shape<1, 1, 1, 64, 1024> v35 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v36 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, pto::Layout::ND> + v37 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + v34 * v12 + v6 * v13), v35, v36); + TLOAD(v32, v37); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + for (size_t v38 = v25; v38 < v24; v38 += v26) { + int64_t v39 = (int64_t)v38; + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v40 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v41 = (uint64_t)v20; + TASSIGN(v40, v41); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + pipe_barrier(PIPE_MTE1); + TEXTRACT(v40, v27, v21, v38); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v42 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v43 = (uint64_t)v21; + TASSIGN(v42, v43); + TEXTRACT(v42, v32, v38, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v44 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v45 = (uint64_t)v19; + TASSIGN(v44, v45); + int64_t v46 = (int64_t)((uint64_t)v39 + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v44, v27, v21, v46); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v47 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v48 = (uint64_t)v18; + TASSIGN(v47, v48); + TEXTRACT(v47, v32, v46, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + if (v39 == v21) { + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v49 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v50 = (uint64_t)v21; + TASSIGN(v49, v50); + pipe_barrier(PIPE_M); + TMATMUL(v49, v40, v42); + } else { + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v52 = (uint64_t)v21; + TASSIGN(v51, v52); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v51, v51, v40, v42); + }; + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v53 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v54 = (uint64_t)v21; + TASSIGN(v53, v54); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + TMATMUL_ACC(v53, v53, v44, v47); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + for (size_t v55 = (size_t)v13; v55 < ((size_t)v9); v55 += (size_t)v8) { + int64_t v56 = (int64_t)((uint64_t)((int64_t)v55) * (uint64_t)v11); + int64_t v57 = (int64_t)((uint64_t)v56 + (uint64_t)v11); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v59 = (uint64_t)v21; + TASSIGN(v58, v59); + pto::Shape<1, 1, 1, 16, 64> v60 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v61 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v62 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v56) * v13), v60, v61 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + TLOAD(v58, v62); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v63 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v64 = (uint64_t)v22; + TASSIGN(v63, v64); + pto::Shape<1, 1, 1, 64, 1024> v65 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v66 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v56) * v12 + v6 * v13), v65, v66); + TLOAD(v63, v67); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v69 = (uint64_t)v17; + TASSIGN(v68, v69); + pto::Shape<1, 1, 1, 16, 64> v70 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v71 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v72 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v57) * v13), v70, v71 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + TLOAD(v68, v72); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v73 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v74 = (uint64_t)v23; + TASSIGN(v73, v74); + pto::Shape<1, 1, 1, 64, 1024> v75 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v76 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND> + v77 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v57) * v12 + v6 * v13), v75, v76); + TLOAD(v73, v77); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + for (size_t v78 = v25; v78 < v24; v78 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v79 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v80 = (uint64_t)v20; + TASSIGN(v79, v80); + pipe_barrier(PIPE_MTE1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + TEXTRACT(v79, v58, v21, v78); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v81 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v82 = (uint64_t)v21; + TASSIGN(v81, v82); + TEXTRACT(v81, v63, v78, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v83 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v84 = (uint64_t)v19; + TASSIGN(v83, v84); + int64_t v85 = (int64_t)((uint64_t)((int64_t)v78) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + TEXTRACT(v83, v58, v21, v85); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v86 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v87 = (uint64_t)v18; + TASSIGN(v86, v87); + TEXTRACT(v86, v63, v85, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v88 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v89 = (uint64_t)v21; + TASSIGN(v88, v89); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v88, v88, v79, v81); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v90 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v91 = (uint64_t)v21; + TASSIGN(v90, v91); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + TMATMUL_ACC(v90, v90, v83, v86); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + }; + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID6); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID6); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + for (size_t v92 = v25; v92 < v24; v92 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v93 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v94 = (uint64_t)v21; + TASSIGN(v93, v94); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + TEXTRACT(v93, v68, v21, v92); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v95 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v96 = (uint64_t)v21; + TASSIGN(v95, v96); + TEXTRACT(v95, v73, v92, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v97 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v98 = (uint64_t)v16; + TASSIGN(v97, v98); + int64_t v99 = (int64_t)((uint64_t)((int64_t)v92) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + TEXTRACT(v97, v68, v21, v99); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v100 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v101 = (uint64_t)v18; + TASSIGN(v100, v101); + TEXTRACT(v100, v73, v99, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v102 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v103 = (uint64_t)v21; + TASSIGN(v102, v103); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v102, v102, v93, v95); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v104 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v105 = (uint64_t)v21; + TASSIGN(v104, v105); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + TMATMUL_ACC(v104, v104, v97, v100); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + }; + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v106 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v107 = (uint64_t)v23; + TASSIGN(v106, v107); + pto::Shape<1, 1, 1, 16, 64> v108 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v109 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v110 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v7) * v13), v108, v109 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + TLOAD(v106, v110); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v111 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v112 = (uint64_t)v22; + TASSIGN(v111, v112); + pto::Shape<1, 1, 1, 64, 1024> v113 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v114 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, pto::Layout::ND> + v115 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v7) * v12 + v6 * v13), v113, v114); + TLOAD(v111, v115); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + for (size_t v116 = v25; v116 < v24; v116 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v117 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v118 = (uint64_t)v20; + TASSIGN(v117, v118); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + TEXTRACT(v117, v106, v21, v116); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v119 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v120 = (uint64_t)v21; + TASSIGN(v119, v120); + TEXTRACT(v119, v111, v116, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v121 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v122 = (uint64_t)v19; + TASSIGN(v121, v122); + int64_t v123 = (int64_t)((uint64_t)((int64_t)v116) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v121, v106, v21, v123); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v124 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v125 = (uint64_t)v18; + TASSIGN(v124, v125); + TEXTRACT(v124, v111, v123, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v126 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v127 = (uint64_t)v21; + TASSIGN(v126, v127); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v126, v126, v117, v119); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v128 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v129 = (uint64_t)v21; + TASSIGN(v128, v129); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + TMATMUL_ACC(v128, v128, v121, v124); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v130 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v131 = (uint64_t)v21; + TASSIGN(v130, v131); + pto::Shape<1, 1, 1, 16, 1024> v132 = pto::Shape<1, 1, 1, 16, 1024>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v133 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v134 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 1024>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v3 + (v21 + v21 * v12), v132, v133 + ); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + TSTORE< + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>, + GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 1024>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>, + AtomicType::AtomicNone>(v134, v130); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); +#endif // __DAV_CUBE__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_norm_in_inline71__rv_v14 + __gm__ Tensor *mlp_norm_in_inline71__rv_v14_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_norm_in_inline71__rv_v14 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_norm_in_inline71__rv_v14_tensor->buffer.addr) + + mlp_norm_in_inline71__rv_v14_tensor->start_offset; + + // Unpack tensor: w_gate__ssa_v0 + __gm__ Tensor *w_gate__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ bfloat16_t *w_gate__ssa_v0 = + reinterpret_cast<__gm__ bfloat16_t *>(w_gate__ssa_v0_tensor->buffer.addr) + w_gate__ssa_v0_tensor->start_offset; + + // Unpack tensor: gate_acc_all_inline203__iter_v11 + __gm__ Tensor *gate_acc_all_inline203__iter_v11_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *gate_acc_all_inline203__iter_v11 = + reinterpret_cast<__gm__ float *>(gate_acc_all_inline203__iter_v11_tensor->buffer.addr) + + gate_acc_all_inline203__iter_v11_tensor->start_offset; + + // Unpack scalar: k0_inline113__ssa_v7 + union { + uint64_t u64; + int64_t val; + } k0_inline113__ssa_v7_conv; + k0_inline113__ssa_v7_conv.u64 = args[3]; + int64_t k0_inline113__ssa_v7 = k0_inline113__ssa_v7_conv.val; + + // Unpack scalar: layer_hidden_base_inline151__ssa_v0 + union { + uint64_t u64; + int64_t val; + } layer_hidden_base_inline151__ssa_v0_conv; + layer_hidden_base_inline151__ssa_v0_conv.u64 = args[4]; + int64_t layer_hidden_base_inline151__ssa_v0 = layer_hidden_base_inline151__ssa_v0_conv.val; + + // Unpack scalar: n0_inline122__ssa_v6 + union { + uint64_t u64; + int64_t val; + } n0_inline122__ssa_v6_conv; + n0_inline122__ssa_v6_conv.u64 = args[5]; + int64_t n0_inline122__ssa_v6 = n0_inline122__ssa_v6_conv.val; + + // Forward to ptoas-generated function + gate_proj_4( + mlp_norm_in_inline71__rv_v14, w_gate__ssa_v0, gate_acc_all_inline203__iter_v11, k0_inline113__ssa_v7, + layer_hidden_base_inline151__ssa_v0, n0_inline122__ssa_v6 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/out_proj.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/out_proj.cpp new file mode 100644 index 0000000000..5f7faa859e --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/out_proj.cpp @@ -0,0 +1,576 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: out_proj + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +static __aicore__ void +out_proj(__gm__ bfloat16_t *v1, __gm__ bfloat16_t *v2, __gm__ float *v3, int64_t v4, int64_t v5, int64_t v6) { + const int64_t v7 = 960; + const int64_t v8 = 2; + const int64_t v9 = 15; + const int64_t v10 = 32; + const int64_t v11 = 512; + const int64_t v12 = 64; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 1024; + const int64_t v17 = 32768; + const int64_t v18 = 3072; + const int64_t v19 = 2048; + const int64_t v20 = 0; + const int64_t v21 = 69632; + const int64_t v22 = 4096; + using T = float; + +#if defined(__DAV_CUBE__) + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v23 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v12); + uint64_t v24 = (uint64_t)v22; + TASSIGN(v23, v24); + pto::Shape<1, 1, 1, 16, 64> v25 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v26 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v27 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v20 + v20 * v14 + v4 * v13), v25, v26 + ); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID4); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + TLOAD(v23, v27); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v28 = Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v12, v11); + uint64_t v29 = (uint64_t)v21; + TASSIGN(v28, v29); + int64_t v30 = (int64_t)((uint64_t)v5 + (uint64_t)v4); + pto::Shape<1, 1, 1, 64, 512> v31 = pto::Shape<1, 1, 1, 64, 512>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v32 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v33 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND>( + v2 + (v20 + v30 * v14 + v6 * v13), v31, v32 + ); + TLOAD(v28, v33); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v34 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v35 = (uint64_t)v19; + TASSIGN(v34, v35); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + TEXTRACT(v34, v23, v20, v20); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v36 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v37 = (uint64_t)v20; + TASSIGN(v36, v37); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v36, v28, v20, v20); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v38 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v39 = (uint64_t)v18; + TASSIGN(v38, v39); + TEXTRACT(v38, v23, v20, v10); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v40 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v41 = (uint64_t)v17; + TASSIGN(v40, v41); + TEXTRACT(v40, v28, v10, v20); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v42 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v43 = (uint64_t)v20; + TASSIGN(v42, v43); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + TMATMUL(v42, v34, v36); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v44 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v45 = (uint64_t)v20; + TASSIGN(v44, v45); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + TMATMUL_ACC(v44, v44, v38, v40); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + for (size_t v46 = (size_t)v13; v46 < ((size_t)v9); v46 += (size_t)v8) { + int64_t v47 = (int64_t)((uint64_t)((int64_t)v46) * (uint64_t)v12); + int64_t v48 = (int64_t)((uint64_t)v47 + (uint64_t)v12); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v49 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v12); + uint64_t v50 = (uint64_t)v20; + TASSIGN(v49, v50); + pto::Shape<1, 1, 1, 16, 64> v51 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v52 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v53 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v20 + v20 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v47) * v13), v51, v52 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + TLOAD(v49, v53); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v54 = Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v12, v11); + uint64_t v55 = (uint64_t)v21; + TASSIGN(v54, v55); + pto::Shape<1, 1, 1, 64, 512> v56 = pto::Shape<1, 1, 1, 64, 512>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v57 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v58 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, + pto::Layout::ND>(v2 + (v20 + (int64_t)((uint64_t)v30 + (uint64_t)v47) * v14 + v6 * v13), v56, v57); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + TLOAD(v54, v58); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v59 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v12); + uint64_t v60 = (uint64_t)v19; + TASSIGN(v59, v60); + pto::Shape<1, 1, 1, 16, 64> v61 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v62 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v63 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v20 + v20 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v48) * v13), v61, v62 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + TLOAD(v59, v63); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID4); + Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v64 = Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v12, v11); + uint64_t v65 = (uint64_t)v22; + TASSIGN(v64, v65); + pto::Shape<1, 1, 1, 64, 512> v66 = pto::Shape<1, 1, 1, 64, 512>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v67 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v68 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, + pto::Layout::ND>(v2 + (v20 + (int64_t)((uint64_t)v30 + (uint64_t)v48) * v14 + v6 * v13), v66, v67); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID4); + TLOAD(v64, v68); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID5); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v69 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v70 = (uint64_t)v19; + TASSIGN(v69, v70); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + TEXTRACT(v69, v49, v20, v20); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v71 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v72 = (uint64_t)v20; + TASSIGN(v71, v72); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v71, v54, v20, v20); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v73 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v74 = (uint64_t)v18; + TASSIGN(v73, v74); + TEXTRACT(v73, v49, v20, v10); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v75 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v76 = (uint64_t)v17; + TASSIGN(v75, v76); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + TEXTRACT(v75, v54, v10, v20); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v77 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v78 = (uint64_t)v20; + TASSIGN(v77, v78); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v77, v77, v69, v71); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v79 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v80 = (uint64_t)v20; + TASSIGN(v79, v80); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + TMATMUL_ACC(v79, v79, v73, v75); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v81 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v82 = (uint64_t)v20; + TASSIGN(v81, v82); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID4); + TEXTRACT(v81, v59, v20, v20); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v83 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v84 = (uint64_t)v20; + TASSIGN(v83, v84); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID5); + TEXTRACT(v83, v64, v20, v20); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v85 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v86 = (uint64_t)v16; + TASSIGN(v85, v86); + TEXTRACT(v85, v59, v20, v10); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v87 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v88 = (uint64_t)v17; + TASSIGN(v87, v88); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + TEXTRACT(v87, v64, v10, v20); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID4); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v89 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v90 = (uint64_t)v20; + TASSIGN(v89, v90); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v89, v89, v81, v83); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v91 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v92 = (uint64_t)v20; + TASSIGN(v91, v92); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + TMATMUL_ACC(v91, v91, v85, v87); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + } + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID5); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v93 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v12); + uint64_t v94 = (uint64_t)v22; + TASSIGN(v93, v94); + pto::Shape<1, 1, 1, 16, 64> v95 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v96 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v97 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v20 + v20 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v7) * v13), v95, v96 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID5); + TLOAD(v93, v97); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID6); + Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v98 = Tile< + TileType::Mat, bfloat16_t, 64, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v12, v11); + uint64_t v99 = (uint64_t)v21; + TASSIGN(v98, v99); + pto::Shape<1, 1, 1, 64, 512> v100 = pto::Shape<1, 1, 1, 64, 512>(); + pto::Stride<327680, 327680, 327680, 5120, 1> v101 = pto::Stride<327680, 327680, 327680, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND> + v102 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 512>, pto::Stride<327680, 327680, 327680, 5120, 1>, pto::Layout::ND>( + v2 + (v20 + (int64_t)((uint64_t)v30 + (uint64_t)v7) * v14 + v6 * v13), v100, v101 + ); + TLOAD(v98, v102); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID7); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v103 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v104 = (uint64_t)v19; + TASSIGN(v103, v104); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID6); + TEXTRACT(v103, v93, v20, v20); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v105 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v106 = (uint64_t)v20; + TASSIGN(v105, v106); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID7); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + TEXTRACT(v105, v98, v20, v20); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v107 = Tile< + TileType::Left, bfloat16_t, 16, 32, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v108 = (uint64_t)v18; + TASSIGN(v107, v108); + TEXTRACT(v107, v93, v20, v10); + Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v109 = Tile< + TileType::Right, bfloat16_t, 32, 512, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null>(v10, v11); + uint64_t v110 = (uint64_t)v17; + TASSIGN(v109, v110); + TEXTRACT(v109, v98, v10, v20); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v111 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v112 = (uint64_t)v20; + TASSIGN(v111, v112); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v111, v111, v103, v105); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v113 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v114 = (uint64_t)v20; + TASSIGN(v113, v114); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + TMATMUL_ACC(v113, v113, v107, v109); + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v115 = Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v116 = (uint64_t)v20; + TASSIGN(v115, v116); + pto::Shape<1, 1, 1, 16, 512> v117 = pto::Shape<1, 1, 1, 16, 512>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v118 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> v119 = + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v20 + v20 * v14), v117, v118 + ); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + TSTORE< + Tile< + TileType::Acc, float, 16, 512, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>, + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>, + AtomicType::AtomicNone>(v119, v115); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID4); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); +#endif // __DAV_CUBE__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: attn_out_inline282__ssa_v4 + __gm__ Tensor *attn_out_inline282__ssa_v4_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *attn_out_inline282__ssa_v4 = + reinterpret_cast<__gm__ bfloat16_t *>(attn_out_inline282__ssa_v4_tensor->buffer.addr) + + attn_out_inline282__ssa_v4_tensor->start_offset; + + // Unpack tensor: wo__ssa_v0 + __gm__ Tensor *wo__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ bfloat16_t *wo__ssa_v0 = + reinterpret_cast<__gm__ bfloat16_t *>(wo__ssa_v0_tensor->buffer.addr) + wo__ssa_v0_tensor->start_offset; + + // Unpack tensor: attn_proj_fp32_inline220__iter_v4 + __gm__ Tensor *attn_proj_fp32_inline220__iter_v4_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *attn_proj_fp32_inline220__iter_v4 = + reinterpret_cast<__gm__ float *>(attn_proj_fp32_inline220__iter_v4_tensor->buffer.addr) + + attn_proj_fp32_inline220__iter_v4_tensor->start_offset; + + // Unpack scalar: k_op_inline266__ssa_v0 + union { + uint64_t u64; + int64_t val; + } k_op_inline266__ssa_v0_conv; + k_op_inline266__ssa_v0_conv.u64 = args[3]; + int64_t k_op_inline266__ssa_v0 = k_op_inline266__ssa_v0_conv.val; + + // Unpack scalar: layer_hidden_base_inline151__ssa_v0 + union { + uint64_t u64; + int64_t val; + } layer_hidden_base_inline151__ssa_v0_conv; + layer_hidden_base_inline151__ssa_v0_conv.u64 = args[4]; + int64_t layer_hidden_base_inline151__ssa_v0 = layer_hidden_base_inline151__ssa_v0_conv.val; + + // Unpack scalar: n_op_inline64__ssa_v0 + union { + uint64_t u64; + int64_t val; + } n_op_inline64__ssa_v0_conv; + n_op_inline64__ssa_v0_conv.u64 = args[5]; + int64_t n_op_inline64__ssa_v0 = n_op_inline64__ssa_v0_conv.val; + + // Forward to ptoas-generated function + out_proj( + attn_out_inline282__ssa_v4, wo__ssa_v0, attn_proj_fp32_inline220__iter_v4, k_op_inline266__ssa_v0, + layer_hidden_base_inline151__ssa_v0, n_op_inline64__ssa_v0 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/up_proj_4.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/up_proj_4.cpp new file mode 100644 index 0000000000..aff7e6b6c4 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aic/up_proj_4.cpp @@ -0,0 +1,621 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: up_proj_4 + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +static __aicore__ void +up_proj_4(__gm__ bfloat16_t *v1, __gm__ bfloat16_t *v2, __gm__ float *v3, int64_t v4, int64_t v5, int64_t v6) { + const int64_t v7 = 960; + const int64_t v8 = 2; + const int64_t v9 = 15; + const int64_t v10 = 32; + const int64_t v11 = 64; + const int64_t v12 = 17408; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 512; + const int64_t v17 = 2048; + const int64_t v18 = 32768; + const int64_t v19 = 1536; + const int64_t v20 = 1024; + const int64_t v21 = 0; + const int64_t v22 = 135168; + const int64_t v23 = 4096; + using T = float; + +#if defined(__DAV_CUBE__) + size_t v24 = (size_t)v11; + size_t v25 = (size_t)v21; + size_t v26 = (size_t)v10; + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v27 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v28 = (uint64_t)v23; + TASSIGN(v27, v28); + pto::Shape<1, 1, 1, 16, 64> v29 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v30 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v31 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + v4 * v13), v29, v30 + ); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + TLOAD(v27, v31); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v32 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v33 = (uint64_t)v22; + TASSIGN(v32, v33); + int64_t v34 = (int64_t)((uint64_t)v5 + (uint64_t)v4); + pto::Shape<1, 1, 1, 64, 1024> v35 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v36 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, pto::Layout::ND> + v37 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + v34 * v12 + v6 * v13), v35, v36); + TLOAD(v32, v37); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + for (size_t v38 = v25; v38 < v24; v38 += v26) { + int64_t v39 = (int64_t)v38; + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v40 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v41 = (uint64_t)v20; + TASSIGN(v40, v41); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + pipe_barrier(PIPE_MTE1); + TEXTRACT(v40, v27, v21, v38); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v42 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v43 = (uint64_t)v21; + TASSIGN(v42, v43); + TEXTRACT(v42, v32, v38, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v44 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v45 = (uint64_t)v19; + TASSIGN(v44, v45); + int64_t v46 = (int64_t)((uint64_t)v39 + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v44, v27, v21, v46); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v47 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v48 = (uint64_t)v18; + TASSIGN(v47, v48); + TEXTRACT(v47, v32, v46, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID0); + if (v39 == v21) { + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v49 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v50 = (uint64_t)v21; + TASSIGN(v49, v50); + pipe_barrier(PIPE_M); + TMATMUL(v49, v40, v42); + } else { + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v52 = (uint64_t)v21; + TASSIGN(v51, v52); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v51, v51, v40, v42); + }; + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v53 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v54 = (uint64_t)v21; + TASSIGN(v53, v54); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID1); + TMATMUL_ACC(v53, v53, v44, v47); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + for (size_t v55 = (size_t)v13; v55 < ((size_t)v9); v55 += (size_t)v8) { + int64_t v56 = (int64_t)((uint64_t)((int64_t)v55) * (uint64_t)v11); + int64_t v57 = (int64_t)((uint64_t)v56 + (uint64_t)v11); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v59 = (uint64_t)v21; + TASSIGN(v58, v59); + pto::Shape<1, 1, 1, 16, 64> v60 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v61 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v62 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v56) * v13), v60, v61 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + TLOAD(v58, v62); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v63 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v64 = (uint64_t)v22; + TASSIGN(v63, v64); + pto::Shape<1, 1, 1, 64, 1024> v65 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v66 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v56) * v12 + v6 * v13), v65, v66); + TLOAD(v63, v67); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v69 = (uint64_t)v17; + TASSIGN(v68, v69); + pto::Shape<1, 1, 1, 16, 64> v70 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v71 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v72 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v57) * v13), v70, v71 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + TLOAD(v68, v72); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v73 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v74 = (uint64_t)v23; + TASSIGN(v73, v74); + pto::Shape<1, 1, 1, 64, 1024> v75 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v76 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND> + v77 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v57) * v12 + v6 * v13), v75, v76); + TLOAD(v73, v77); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + for (size_t v78 = v25; v78 < v24; v78 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v79 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v80 = (uint64_t)v20; + TASSIGN(v79, v80); + pipe_barrier(PIPE_MTE1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + TEXTRACT(v79, v58, v21, v78); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v81 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v82 = (uint64_t)v21; + TASSIGN(v81, v82); + TEXTRACT(v81, v63, v78, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v83 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v84 = (uint64_t)v19; + TASSIGN(v83, v84); + int64_t v85 = (int64_t)((uint64_t)((int64_t)v78) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + TEXTRACT(v83, v58, v21, v85); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v86 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v87 = (uint64_t)v18; + TASSIGN(v86, v87); + TEXTRACT(v86, v63, v85, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v88 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v89 = (uint64_t)v21; + TASSIGN(v88, v89); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID2); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v88, v88, v79, v81); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v90 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v91 = (uint64_t)v21; + TASSIGN(v90, v91); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID3); + TMATMUL_ACC(v90, v90, v83, v86); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + }; + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID4); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID5); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID6); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID6); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + for (size_t v92 = v25; v92 < v24; v92 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v93 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v94 = (uint64_t)v21; + TASSIGN(v93, v94); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + TEXTRACT(v93, v68, v21, v92); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v95 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v96 = (uint64_t)v21; + TASSIGN(v95, v96); + TEXTRACT(v95, v73, v92, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v97 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v15); + uint64_t v98 = (uint64_t)v16; + TASSIGN(v97, v98); + int64_t v99 = (int64_t)((uint64_t)((int64_t)v92) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + TEXTRACT(v97, v68, v21, v99); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null> + v100 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v101 = (uint64_t)v18; + TASSIGN(v100, v101); + TEXTRACT(v100, v73, v99, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v102 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v103 = (uint64_t)v21; + TASSIGN(v102, v103); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID4); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v102, v102, v93, v95); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v104 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v105 = (uint64_t)v21; + TASSIGN(v104, v105); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID5); + TMATMUL_ACC(v104, v104, v97, v100); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + }; + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID7); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID3); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v106 = Tile< + TileType::Mat, bfloat16_t, 16, 64, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v11); + uint64_t v107 = (uint64_t)v23; + TASSIGN(v106, v107); + pto::Shape<1, 1, 1, 16, 64> v108 = pto::Shape<1, 1, 1, 16, 64>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v109 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v110 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 64>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v21 + v21 * v14 + (int64_t)((uint64_t)v4 + (uint64_t)v7) * v13), v108, v109 + ); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID3); + TLOAD(v106, v110); + Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v111 = Tile< + TileType::Mat, bfloat16_t, 64, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v11, v20); + uint64_t v112 = (uint64_t)v22; + TASSIGN(v111, v112); + pto::Shape<1, 1, 1, 64, 1024> v113 = pto::Shape<1, 1, 1, 64, 1024>(); + pto::Stride<1114112, 1114112, 1114112, 17408, 1> v114 = pto::Stride<1114112, 1114112, 1114112, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, pto::Layout::ND> + v115 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 64, 1024>, pto::Stride<1114112, 1114112, 1114112, 17408, 1>, + pto::Layout::ND>(v2 + (v21 + (int64_t)((uint64_t)v34 + (uint64_t)v7) * v12 + v6 * v13), v113, v114); + TLOAD(v111, v115); + set_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_MTE2, PIPE_MTE1, EVENT_ID3); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + for (size_t v116 = v25; v116 < v24; v116 += v26) { + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v117 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v118 = (uint64_t)v20; + TASSIGN(v117, v118); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + TEXTRACT(v117, v106, v21, v116); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v119 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v120 = (uint64_t)v21; + TASSIGN(v119, v120); + TEXTRACT(v119, v111, v116, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null> + v121 = Tile< + TileType::Left, bfloat16_t, 16, 16, BLayout::RowMajor, -1, -1, SLayout::RowMajor, 512, PadValue::Null, + CompactMode::Null>(v15, v15); + uint64_t v122 = (uint64_t)v19; + TASSIGN(v121, v122); + int64_t v123 = (int64_t)((uint64_t)((int64_t)v116) + (uint64_t)v15); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + TEXTRACT(v121, v106, v21, v123); + Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, PadValue::Null, + CompactMode::Null> + v124 = Tile< + TileType::Right, bfloat16_t, 16, 1024, BLayout::RowMajor, -1, -1, SLayout::ColMajor, 512, + PadValue::Null, CompactMode::Null>(v15, v20); + uint64_t v125 = (uint64_t)v18; + TASSIGN(v124, v125); + TEXTRACT(v124, v111, v123, v21); + set_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v126 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v127 = (uint64_t)v21; + TASSIGN(v126, v127); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID6); + pipe_barrier(PIPE_M); + TMATMUL_ACC(v126, v126, v117, v119); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v128 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v129 = (uint64_t)v21; + TASSIGN(v128, v129); + pipe_barrier(PIPE_M); + wait_flag(PIPE_MTE1, PIPE_M, EVENT_ID7); + TMATMUL_ACC(v128, v128, v121, v124); + set_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + } + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID1); + wait_flag(PIPE_M, PIPE_MTE1, EVENT_ID0); + set_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null> + v130 = Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>(v15, v20); + uint64_t v131 = (uint64_t)v21; + TASSIGN(v130, v131); + pto::Shape<1, 1, 1, 16, 1024> v132 = pto::Shape<1, 1, 1, 16, 1024>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v133 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v134 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 1024>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v3 + (v21 + v21 * v12), v132, v133 + ); + wait_flag(PIPE_M, PIPE_FIX, EVENT_ID0); + TSTORE< + Tile< + TileType::Acc, float, 16, 1024, BLayout::ColMajor, -1, -1, SLayout::RowMajor, 1024, PadValue::Null, + CompactMode::Null>, + GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 1024>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>, + AtomicType::AtomicNone>(v134, v130); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE1, PIPE_MTE2, EVENT_ID2); +#endif // __DAV_CUBE__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_norm_in_inline71__rv_v14 + __gm__ Tensor *mlp_norm_in_inline71__rv_v14_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_norm_in_inline71__rv_v14 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_norm_in_inline71__rv_v14_tensor->buffer.addr) + + mlp_norm_in_inline71__rv_v14_tensor->start_offset; + + // Unpack tensor: w_up__ssa_v0 + __gm__ Tensor *w_up__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ bfloat16_t *w_up__ssa_v0 = + reinterpret_cast<__gm__ bfloat16_t *>(w_up__ssa_v0_tensor->buffer.addr) + w_up__ssa_v0_tensor->start_offset; + + // Unpack tensor: up_acc_all_inline303__iter_v11 + __gm__ Tensor *up_acc_all_inline303__iter_v11_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *up_acc_all_inline303__iter_v11 = + reinterpret_cast<__gm__ float *>(up_acc_all_inline303__iter_v11_tensor->buffer.addr) + + up_acc_all_inline303__iter_v11_tensor->start_offset; + + // Unpack scalar: k0_inline113__ssa_v7 + union { + uint64_t u64; + int64_t val; + } k0_inline113__ssa_v7_conv; + k0_inline113__ssa_v7_conv.u64 = args[3]; + int64_t k0_inline113__ssa_v7 = k0_inline113__ssa_v7_conv.val; + + // Unpack scalar: layer_hidden_base_inline151__ssa_v0 + union { + uint64_t u64; + int64_t val; + } layer_hidden_base_inline151__ssa_v0_conv; + layer_hidden_base_inline151__ssa_v0_conv.u64 = args[4]; + int64_t layer_hidden_base_inline151__ssa_v0 = layer_hidden_base_inline151__ssa_v0_conv.val; + + // Unpack scalar: n0_inline122__ssa_v6 + union { + uint64_t u64; + int64_t val; + } n0_inline122__ssa_v6_conv; + n0_inline122__ssa_v6_conv.u64 = args[5]; + int64_t n0_inline122__ssa_v6 = n0_inline122__ssa_v6_conv.val; + + // Forward to ptoas-generated function + up_proj_4( + mlp_norm_in_inline71__rv_v14, w_up__ssa_v0, up_acc_all_inline303__iter_v11, k0_inline113__ssa_v7, + layer_hidden_base_inline151__ssa_v0, n0_inline122__ssa_v6 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce.h b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce.h new file mode 100644 index 0000000000..862f5041d1 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce.h @@ -0,0 +1,70 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +#pragma once + +#include + +#include + +// Sums the private split-K partials of one output band. +// +// `part` addresses row 0 of a [splits*16, RowStride] fp32 partial buffer, already +// offset to the band's first column; partial j occupies rows [j*16, j*16+16), so +// its element offset from `part` is j * 16 * RowStride. `out` addresses the same +// band of the [16, RowStride] accumulator. Both carry row stride RowStride +// because each is a column slice of its parent, so a row step is a full parent row. +// +// `Atomic` is AtomicNone when the reducer is the band's only writer, and +// AtomicAdd when another task still accumulates into the same band and the sum +// has to join it rather than replace it. +// +// `band_width` must be a multiple of 256; `splits` must be at least 1. +template +static __aicore__ inline void +reduce_partials_into_band(__gm__ float *part, __gm__ float *out, int64_t splits, int64_t band_width) { + constexpr int64_t kRows = 16; + constexpr int64_t kChunk = 256; + constexpr int64_t kPartialStride = kRows * RowStride; + constexpr uint64_t kAccUb = 0; + constexpr uint64_t kTmpUb = kRows * kChunk * sizeof(float); + + using VecTile = pto::Tile< + pto::TileType::Vec, float, kRows, kChunk, pto::BLayout::RowMajor, -1, -1, pto::SLayout::NoneBox, 512, + pto::PadValue::Null, pto::CompactMode::Null>; + using ChunkShape = pto::Shape<1, 1, 1, kRows, kChunk>; + using ChunkStride = pto::Stride; + using ChunkTensor = pto::GlobalTensor; + + ChunkShape shape = ChunkShape(); + ChunkStride stride = ChunkStride(); + + for (int64_t c = 0; c < band_width; c += kChunk) { + VecTile acc = VecTile(kRows, kChunk); + TASSIGN(acc, kAccUb); + ChunkTensor first = ChunkTensor(part + c, shape, stride); + TLOAD(acc, first); + pipe_barrier(PIPE_ALL); + + for (int64_t j = 1; j < splits; ++j) { + VecTile tmp = VecTile(kRows, kChunk); + TASSIGN(tmp, kTmpUb); + ChunkTensor src = ChunkTensor(part + j * kPartialStride + c, shape, stride); + TLOAD(tmp, src); + pipe_barrier(PIPE_ALL); + TADD(acc, acc, tmp); + pipe_barrier(PIPE_ALL); + } + + ChunkTensor dst = ChunkTensor(out + c, shape, stride); + TSTORE(dst, acc); + pipe_barrier(PIPE_ALL); + } +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden.cpp new file mode 100644 index 0000000000..dc9279095e --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden.cpp @@ -0,0 +1,69 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: partials_reduce_hidden +// +// Reducer over a hidden-width accumulator (row stride 5120). Used for +// down_acc_all, whose bands have no writer other than the split-K partials this +// sums, so the result is stored rather than accumulated. + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +#include "partials_reduce.h" + +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *part_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ float *part = reinterpret_cast<__gm__ float *>(part_tensor->buffer.addr) + part_tensor->start_offset; + + __gm__ Tensor *out_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *out = reinterpret_cast<__gm__ float *>(out_tensor->buffer.addr) + out_tensor->start_offset; + + union { + uint64_t u64; + int64_t val; + } splits_conv; + splits_conv.u64 = args[2]; + int64_t splits = splits_conv.val; + + union { + uint64_t u64; + int64_t val; + } band_width_conv; + band_width_conv.u64 = args[3]; + int64_t band_width = band_width_conv.val; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + reduce_partials_into_band<5120, pto::AtomicType::AtomicNone>(part, out, splits, band_width); +#else + (void)part; + (void)out; + (void)splits; + (void)band_width; +#endif // __DAV_VEC__ + + pipe_barrier(PIPE_ALL); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden_add.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden_add.cpp new file mode 100644 index 0000000000..3673aa8f65 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_hidden_add.cpp @@ -0,0 +1,70 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: partials_reduce_hidden_add +// +// Reducer over a hidden-width accumulator (row stride 5120) whose bands have a +// second writer. Used for attn_proj_fp32: out_proj_0's SPMD blocks accumulate +// into the same bands from a block_idx-derived offset that orchestration cannot +// name, so this sum joins them atomically rather than replacing them. + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +#include "partials_reduce.h" + +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *part_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ float *part = reinterpret_cast<__gm__ float *>(part_tensor->buffer.addr) + part_tensor->start_offset; + + __gm__ Tensor *out_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *out = reinterpret_cast<__gm__ float *>(out_tensor->buffer.addr) + out_tensor->start_offset; + + union { + uint64_t u64; + int64_t val; + } splits_conv; + splits_conv.u64 = args[2]; + int64_t splits = splits_conv.val; + + union { + uint64_t u64; + int64_t val; + } band_width_conv; + band_width_conv.u64 = args[3]; + int64_t band_width = band_width_conv.val; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + reduce_partials_into_band<5120, pto::AtomicType::AtomicAdd>(part, out, splits, band_width); +#else + (void)part; + (void)out; + (void)splits; + (void)band_width; +#endif // __DAV_VEC__ + + pipe_barrier(PIPE_ALL); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_inter.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_inter.cpp new file mode 100644 index 0000000000..03bc84be58 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/partials_reduce_inter.cpp @@ -0,0 +1,69 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: partials_reduce_inter +// +// Reducer over an intermediate-width accumulator (row stride 17408). Used for +// gate_acc_all and up_acc_all, whose bands have no writer other than the split-K +// partials this sums, so the result is stored rather than accumulated. + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +#include "partials_reduce.h" + +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + __gm__ Tensor *part_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ float *part = reinterpret_cast<__gm__ float *>(part_tensor->buffer.addr) + part_tensor->start_offset; + + __gm__ Tensor *out_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *out = reinterpret_cast<__gm__ float *>(out_tensor->buffer.addr) + out_tensor->start_offset; + + union { + uint64_t u64; + int64_t val; + } splits_conv; + splits_conv.u64 = args[2]; + int64_t splits = splits_conv.val; + + union { + uint64_t u64; + int64_t val; + } band_width_conv; + band_width_conv.u64 = args[3]; + int64_t band_width = band_width_conv.val; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + reduce_partials_into_band<17408, pto::AtomicType::AtomicNone>(part, out, splits, band_width); +#else + (void)part; + (void)out; + (void)splits; + (void)band_width; +#endif // __DAV_VEC__ + + pipe_barrier(PIPE_ALL); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast.cpp new file mode 100644 index 0000000000..cb3239ad56 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast.cpp @@ -0,0 +1,369 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: residual_rms_cast + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +// The two stores below address the band relatively: the task's `mlp_norm_in` and +// `post_norm_partial` arguments are views narrowed to [k_base, k_base + 1024), so +// their base already carries k_base and only the in-band offset is added. The +// loads keep the parent views and so still index absolutely. +static __aicore__ void residual_rms_cast( + __gm__ bfloat16_t *v1, __gm__ float *v2, __gm__ float *v3, __gm__ float *v4, __gm__ float *v5, int64_t v6, + int64_t v7 +) { + SaturationMode v8 = SaturationMode::OFF; + RoundMode v9 = RoundMode::CAST_ROUND; + const int64_t v10 = 256; + const int64_t v11 = 2; + const int64_t v12 = 4; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 0; + const int64_t v17 = 51200; + const int64_t v18 = 34816; + const int64_t v19 = 33792; + const int64_t v20 = 17408; + const int64_t v21 = 1024; + using T = float; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + for (size_t v22 = (size_t)v16; v22 < ((size_t)v12); v22 += (size_t)v11) { + int64_t v23 = (int64_t)((uint64_t)((int64_t)v22) * (uint64_t)v10); + int64_t v24 = (int64_t)((uint64_t)v6 + (uint64_t)v23); + int64_t v25 = (int64_t)((uint64_t)v6 + (uint64_t)((int64_t)(uint64_t)v23 + (uint64_t)v10)); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v26 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v27 = (uint64_t)v21; + TASSIGN(v26, v27); + pto::Shape<1, 1, 1, 16, 256> v28 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v29 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v30 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v24 * v13), v28, v29 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + TLOAD(v26, v30); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v31 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v32 = (uint64_t)v20; + TASSIGN(v31, v32); + pto::Shape<1, 1, 1, 16, 256> v33 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v34 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v35 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v24 * v13), v33, v34 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + TLOAD(v31, v35); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v36 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v37 = (uint64_t)v19; + TASSIGN(v36, v37); + pto::Shape<1, 1, 1, 1, 256> v38 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v39 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v40 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v24 * v13), v38, v39 + ); + TLOAD(v36, v40); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v41 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v42 = (uint64_t)v18; + TASSIGN(v41, v42); + pto::Shape<1, 1, 1, 16, 256> v43 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v44 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v45 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v25 * v13), v43, v44 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + TLOAD(v41, v45); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v46 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v47 = (uint64_t)v17; + TASSIGN(v46, v47); + pto::Shape<1, 1, 1, 16, 256> v48 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v49 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v50 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v25 * v13), v48, v49 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + TLOAD(v46, v50); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v52 = (uint64_t)v16; + TASSIGN(v51, v52); + pto::Shape<1, 1, 1, 1, 256> v53 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v54 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v55 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v25 * v13), v53, v54 + ); + TLOAD(v51, v55); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v56 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v57 = (uint64_t)v21; + TASSIGN(v56, v57); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + TADD(v56, v26, v31); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v59 = (uint64_t)v20; + TASSIGN(v58, v59); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TCOLEXPANDMUL(v58, v56, v36); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v60 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v61 = (uint64_t)v20; + TASSIGN(v60, v61); + pipe_barrier(PIPE_V); + TCVT(v60, v58, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + pto::Shape<1, 1, 1, 16, 256> v62 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v63 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v64 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + v23 * v13), v62, v63 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + pipe_barrier(PIPE_MTE3); + TSTORE(v64, v56); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + pto::Shape<1, 1, 1, 16, 256> v65 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v66 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + v23 * v13), v65, v66 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(v67, v60); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v69 = (uint64_t)v18; + TASSIGN(v68, v69); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + TADD(v68, v41, v46); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v70 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v71 = (uint64_t)v17; + TASSIGN(v70, v71); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + TCOLEXPANDMUL(v70, v68, v51); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v72 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v73 = (uint64_t)v17; + TASSIGN(v72, v73); + pipe_barrier(PIPE_V); + TCVT(v72, v70, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + pto::Shape<1, 1, 1, 16, 256> v74 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v75 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v76 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + (v23 + v10) * v13), v74, v75 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + pipe_barrier(PIPE_MTE3); + TSTORE(v76, v68); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + pto::Shape<1, 1, 1, 16, 256> v77 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v78 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v79 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + (v23 + v10) * v13), v77, v78 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + TSTORE(v79, v72); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + } + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); +#endif // __DAV_VEC__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_norm_in_inline71__ssa_v0 + __gm__ Tensor *mlp_norm_in_inline71__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_norm_in_inline71__ssa_v0 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_norm_in_inline71__ssa_v0_tensor->buffer.addr) + + mlp_norm_in_inline71__ssa_v0_tensor->start_offset; + + // Unpack tensor: post_norm_partial_inline118__ssa_v0 + __gm__ Tensor *post_norm_partial_inline118__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *post_norm_partial_inline118__ssa_v0 = + reinterpret_cast<__gm__ float *>(post_norm_partial_inline118__ssa_v0_tensor->buffer.addr) + + post_norm_partial_inline118__ssa_v0_tensor->start_offset; + + // Unpack tensor: attn_proj_fp32_inline220__ssa_v7 + __gm__ Tensor *attn_proj_fp32_inline220__ssa_v7_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *attn_proj_fp32_inline220__ssa_v7 = + reinterpret_cast<__gm__ float *>(attn_proj_fp32_inline220__ssa_v7_tensor->buffer.addr) + + attn_proj_fp32_inline220__ssa_v7_tensor->start_offset; + + // Unpack tensor: cur__iter_v6 + __gm__ Tensor *cur__iter_v6_tensor = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ float *cur__iter_v6 = + reinterpret_cast<__gm__ float *>(cur__iter_v6_tensor->buffer.addr) + cur__iter_v6_tensor->start_offset; + + // Unpack tensor: post_rms_weight__ssa_v0 + __gm__ Tensor *post_rms_weight__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ float *post_rms_weight__ssa_v0 = + reinterpret_cast<__gm__ float *>(post_rms_weight__ssa_v0_tensor->buffer.addr) + + post_rms_weight__ssa_v0_tensor->start_offset; + + // Unpack scalar: k_base_inline111__ssa_v0 + union { + uint64_t u64; + int64_t val; + } k_base_inline111__ssa_v0_conv; + k_base_inline111__ssa_v0_conv.u64 = args[5]; + int64_t k_base_inline111__ssa_v0 = k_base_inline111__ssa_v0_conv.val; + + // Unpack scalar: i__idx_v0 + union { + uint64_t u64; + int64_t val; + } i__idx_v0_conv; + i__idx_v0_conv.u64 = args[6]; + int64_t i__idx_v0 = i__idx_v0_conv.val; + + // Forward to ptoas-generated function + residual_rms_cast( + mlp_norm_in_inline71__ssa_v0, post_norm_partial_inline118__ssa_v0, attn_proj_fp32_inline220__ssa_v7, + cur__iter_v6, post_rms_weight__ssa_v0, k_base_inline111__ssa_v0, i__idx_v0 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_0.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_0.cpp new file mode 100644 index 0000000000..5d550faf8e --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_0.cpp @@ -0,0 +1,369 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: residual_rms_cast_0 + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +// The two stores below address the band relatively: the task's `mlp_norm_in` and +// `post_norm_partial` arguments are views narrowed to [k_base, k_base + 1024), so +// their base already carries k_base and only the in-band offset is added. The +// loads keep the parent views and so still index absolutely. +static __aicore__ void residual_rms_cast_0( + __gm__ bfloat16_t *v1, __gm__ float *v2, __gm__ float *v3, __gm__ float *v4, __gm__ float *v5, int64_t v6, + int64_t v7 +) { + SaturationMode v8 = SaturationMode::OFF; + RoundMode v9 = RoundMode::CAST_ROUND; + const int64_t v10 = 256; + const int64_t v11 = 2; + const int64_t v12 = 4; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 0; + const int64_t v17 = 51200; + const int64_t v18 = 34816; + const int64_t v19 = 33792; + const int64_t v20 = 17408; + const int64_t v21 = 1024; + using T = float; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + for (size_t v22 = (size_t)v16; v22 < ((size_t)v12); v22 += (size_t)v11) { + int64_t v23 = (int64_t)((uint64_t)((int64_t)v22) * (uint64_t)v10); + int64_t v24 = (int64_t)((uint64_t)v6 + (uint64_t)v23); + int64_t v25 = (int64_t)((uint64_t)v6 + (uint64_t)((int64_t)(uint64_t)v23 + (uint64_t)v10)); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v26 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v27 = (uint64_t)v21; + TASSIGN(v26, v27); + pto::Shape<1, 1, 1, 16, 256> v28 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v29 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v30 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v24 * v13), v28, v29 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + TLOAD(v26, v30); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v31 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v32 = (uint64_t)v20; + TASSIGN(v31, v32); + pto::Shape<1, 1, 1, 16, 256> v33 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v34 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v35 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v24 * v13), v33, v34 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + TLOAD(v31, v35); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v36 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v37 = (uint64_t)v19; + TASSIGN(v36, v37); + pto::Shape<1, 1, 1, 1, 256> v38 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v39 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v40 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v24 * v13), v38, v39 + ); + TLOAD(v36, v40); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v41 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v42 = (uint64_t)v18; + TASSIGN(v41, v42); + pto::Shape<1, 1, 1, 16, 256> v43 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v44 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v45 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v25 * v13), v43, v44 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + TLOAD(v41, v45); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v46 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v47 = (uint64_t)v17; + TASSIGN(v46, v47); + pto::Shape<1, 1, 1, 16, 256> v48 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v49 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v50 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v25 * v13), v48, v49 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + TLOAD(v46, v50); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v52 = (uint64_t)v16; + TASSIGN(v51, v52); + pto::Shape<1, 1, 1, 1, 256> v53 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v54 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v55 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v25 * v13), v53, v54 + ); + TLOAD(v51, v55); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v56 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v57 = (uint64_t)v21; + TASSIGN(v56, v57); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + TADD(v56, v26, v31); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v59 = (uint64_t)v20; + TASSIGN(v58, v59); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TCOLEXPANDMUL(v58, v56, v36); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v60 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v61 = (uint64_t)v20; + TASSIGN(v60, v61); + pipe_barrier(PIPE_V); + TCVT(v60, v58, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + pto::Shape<1, 1, 1, 16, 256> v62 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v63 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v64 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + v23 * v13), v62, v63 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + pipe_barrier(PIPE_MTE3); + TSTORE(v64, v56); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + pto::Shape<1, 1, 1, 16, 256> v65 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v66 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + v23 * v13), v65, v66 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(v67, v60); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v69 = (uint64_t)v18; + TASSIGN(v68, v69); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + TADD(v68, v41, v46); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v70 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v71 = (uint64_t)v17; + TASSIGN(v70, v71); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + TCOLEXPANDMUL(v70, v68, v51); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v72 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v73 = (uint64_t)v17; + TASSIGN(v72, v73); + pipe_barrier(PIPE_V); + TCVT(v72, v70, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + pto::Shape<1, 1, 1, 16, 256> v74 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v75 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v76 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + (v23 + v10) * v13), v74, v75 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + pipe_barrier(PIPE_MTE3); + TSTORE(v76, v68); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + pto::Shape<1, 1, 1, 16, 256> v77 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v78 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v79 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + (v23 + v10) * v13), v77, v78 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + TSTORE(v79, v72); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + } + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); +#endif // __DAV_VEC__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_norm_in_inline71__rv_v2 + __gm__ Tensor *mlp_norm_in_inline71__rv_v2_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_norm_in_inline71__rv_v2 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_norm_in_inline71__rv_v2_tensor->buffer.addr) + + mlp_norm_in_inline71__rv_v2_tensor->start_offset; + + // Unpack tensor: post_norm_partial_inline118__rv_v2 + __gm__ Tensor *post_norm_partial_inline118__rv_v2_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *post_norm_partial_inline118__rv_v2 = + reinterpret_cast<__gm__ float *>(post_norm_partial_inline118__rv_v2_tensor->buffer.addr) + + post_norm_partial_inline118__rv_v2_tensor->start_offset; + + // Unpack tensor: attn_proj_fp32_inline220__ssa_v7 + __gm__ Tensor *attn_proj_fp32_inline220__ssa_v7_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *attn_proj_fp32_inline220__ssa_v7 = + reinterpret_cast<__gm__ float *>(attn_proj_fp32_inline220__ssa_v7_tensor->buffer.addr) + + attn_proj_fp32_inline220__ssa_v7_tensor->start_offset; + + // Unpack tensor: cur__iter_v6 + __gm__ Tensor *cur__iter_v6_tensor = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ float *cur__iter_v6 = + reinterpret_cast<__gm__ float *>(cur__iter_v6_tensor->buffer.addr) + cur__iter_v6_tensor->start_offset; + + // Unpack tensor: post_rms_weight__ssa_v0 + __gm__ Tensor *post_rms_weight__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ float *post_rms_weight__ssa_v0 = + reinterpret_cast<__gm__ float *>(post_rms_weight__ssa_v0_tensor->buffer.addr) + + post_rms_weight__ssa_v0_tensor->start_offset; + + // Unpack scalar: k_base_inline111__ssa_v1 + union { + uint64_t u64; + int64_t val; + } k_base_inline111__ssa_v1_conv; + k_base_inline111__ssa_v1_conv.u64 = args[5]; + int64_t k_base_inline111__ssa_v1 = k_base_inline111__ssa_v1_conv.val; + + // Unpack scalar: i__idx_v0 + union { + uint64_t u64; + int64_t val; + } i__idx_v0_conv; + i__idx_v0_conv.u64 = args[6]; + int64_t i__idx_v0 = i__idx_v0_conv.val; + + // Forward to ptoas-generated function + residual_rms_cast_0( + mlp_norm_in_inline71__rv_v2, post_norm_partial_inline118__rv_v2, attn_proj_fp32_inline220__ssa_v7, cur__iter_v6, + post_rms_weight__ssa_v0, k_base_inline111__ssa_v1, i__idx_v0 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_1.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_1.cpp new file mode 100644 index 0000000000..b7b3188e4e --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_1.cpp @@ -0,0 +1,369 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: residual_rms_cast_1 + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +// The two stores below address the band relatively: the task's `mlp_norm_in` and +// `post_norm_partial` arguments are views narrowed to [k_base, k_base + 1024), so +// their base already carries k_base and only the in-band offset is added. The +// loads keep the parent views and so still index absolutely. +static __aicore__ void residual_rms_cast_1( + __gm__ bfloat16_t *v1, __gm__ float *v2, __gm__ float *v3, __gm__ float *v4, __gm__ float *v5, int64_t v6, + int64_t v7 +) { + SaturationMode v8 = SaturationMode::OFF; + RoundMode v9 = RoundMode::CAST_ROUND; + const int64_t v10 = 256; + const int64_t v11 = 2; + const int64_t v12 = 4; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 0; + const int64_t v17 = 51200; + const int64_t v18 = 34816; + const int64_t v19 = 33792; + const int64_t v20 = 17408; + const int64_t v21 = 1024; + using T = float; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + for (size_t v22 = (size_t)v16; v22 < ((size_t)v12); v22 += (size_t)v11) { + int64_t v23 = (int64_t)((uint64_t)((int64_t)v22) * (uint64_t)v10); + int64_t v24 = (int64_t)((uint64_t)v6 + (uint64_t)v23); + int64_t v25 = (int64_t)((uint64_t)v6 + (uint64_t)((int64_t)(uint64_t)v23 + (uint64_t)v10)); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v26 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v27 = (uint64_t)v21; + TASSIGN(v26, v27); + pto::Shape<1, 1, 1, 16, 256> v28 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v29 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v30 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v24 * v13), v28, v29 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + TLOAD(v26, v30); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v31 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v32 = (uint64_t)v20; + TASSIGN(v31, v32); + pto::Shape<1, 1, 1, 16, 256> v33 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v34 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v35 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v24 * v13), v33, v34 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + TLOAD(v31, v35); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v36 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v37 = (uint64_t)v19; + TASSIGN(v36, v37); + pto::Shape<1, 1, 1, 1, 256> v38 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v39 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v40 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v24 * v13), v38, v39 + ); + TLOAD(v36, v40); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v41 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v42 = (uint64_t)v18; + TASSIGN(v41, v42); + pto::Shape<1, 1, 1, 16, 256> v43 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v44 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v45 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v25 * v13), v43, v44 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + TLOAD(v41, v45); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v46 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v47 = (uint64_t)v17; + TASSIGN(v46, v47); + pto::Shape<1, 1, 1, 16, 256> v48 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v49 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v50 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v25 * v13), v48, v49 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + TLOAD(v46, v50); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v52 = (uint64_t)v16; + TASSIGN(v51, v52); + pto::Shape<1, 1, 1, 1, 256> v53 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v54 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v55 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v25 * v13), v53, v54 + ); + TLOAD(v51, v55); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v56 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v57 = (uint64_t)v21; + TASSIGN(v56, v57); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + TADD(v56, v26, v31); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v59 = (uint64_t)v20; + TASSIGN(v58, v59); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TCOLEXPANDMUL(v58, v56, v36); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v60 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v61 = (uint64_t)v20; + TASSIGN(v60, v61); + pipe_barrier(PIPE_V); + TCVT(v60, v58, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + pto::Shape<1, 1, 1, 16, 256> v62 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v63 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v64 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + v23 * v13), v62, v63 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + pipe_barrier(PIPE_MTE3); + TSTORE(v64, v56); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + pto::Shape<1, 1, 1, 16, 256> v65 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v66 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + v23 * v13), v65, v66 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(v67, v60); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v69 = (uint64_t)v18; + TASSIGN(v68, v69); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + TADD(v68, v41, v46); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v70 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v71 = (uint64_t)v17; + TASSIGN(v70, v71); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + TCOLEXPANDMUL(v70, v68, v51); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v72 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v73 = (uint64_t)v17; + TASSIGN(v72, v73); + pipe_barrier(PIPE_V); + TCVT(v72, v70, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + pto::Shape<1, 1, 1, 16, 256> v74 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v75 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v76 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + (v23 + v10) * v13), v74, v75 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + pipe_barrier(PIPE_MTE3); + TSTORE(v76, v68); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + pto::Shape<1, 1, 1, 16, 256> v77 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v78 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v79 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + (v23 + v10) * v13), v77, v78 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + TSTORE(v79, v72); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + } + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); +#endif // __DAV_VEC__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_norm_in_inline71__rv_v5 + __gm__ Tensor *mlp_norm_in_inline71__rv_v5_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_norm_in_inline71__rv_v5 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_norm_in_inline71__rv_v5_tensor->buffer.addr) + + mlp_norm_in_inline71__rv_v5_tensor->start_offset; + + // Unpack tensor: post_norm_partial_inline118__rv_v5 + __gm__ Tensor *post_norm_partial_inline118__rv_v5_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *post_norm_partial_inline118__rv_v5 = + reinterpret_cast<__gm__ float *>(post_norm_partial_inline118__rv_v5_tensor->buffer.addr) + + post_norm_partial_inline118__rv_v5_tensor->start_offset; + + // Unpack tensor: attn_proj_fp32_inline220__ssa_v7 + __gm__ Tensor *attn_proj_fp32_inline220__ssa_v7_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *attn_proj_fp32_inline220__ssa_v7 = + reinterpret_cast<__gm__ float *>(attn_proj_fp32_inline220__ssa_v7_tensor->buffer.addr) + + attn_proj_fp32_inline220__ssa_v7_tensor->start_offset; + + // Unpack tensor: cur__iter_v6 + __gm__ Tensor *cur__iter_v6_tensor = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ float *cur__iter_v6 = + reinterpret_cast<__gm__ float *>(cur__iter_v6_tensor->buffer.addr) + cur__iter_v6_tensor->start_offset; + + // Unpack tensor: post_rms_weight__ssa_v0 + __gm__ Tensor *post_rms_weight__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ float *post_rms_weight__ssa_v0 = + reinterpret_cast<__gm__ float *>(post_rms_weight__ssa_v0_tensor->buffer.addr) + + post_rms_weight__ssa_v0_tensor->start_offset; + + // Unpack scalar: k_base_inline111__ssa_v2 + union { + uint64_t u64; + int64_t val; + } k_base_inline111__ssa_v2_conv; + k_base_inline111__ssa_v2_conv.u64 = args[5]; + int64_t k_base_inline111__ssa_v2 = k_base_inline111__ssa_v2_conv.val; + + // Unpack scalar: i__idx_v0 + union { + uint64_t u64; + int64_t val; + } i__idx_v0_conv; + i__idx_v0_conv.u64 = args[6]; + int64_t i__idx_v0 = i__idx_v0_conv.val; + + // Forward to ptoas-generated function + residual_rms_cast_1( + mlp_norm_in_inline71__rv_v5, post_norm_partial_inline118__rv_v5, attn_proj_fp32_inline220__ssa_v7, cur__iter_v6, + post_rms_weight__ssa_v0, k_base_inline111__ssa_v2, i__idx_v0 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_2.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_2.cpp new file mode 100644 index 0000000000..c397839bed --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_2.cpp @@ -0,0 +1,369 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: residual_rms_cast_2 + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +// The two stores below address the band relatively: the task's `mlp_norm_in` and +// `post_norm_partial` arguments are views narrowed to [k_base, k_base + 1024), so +// their base already carries k_base and only the in-band offset is added. The +// loads keep the parent views and so still index absolutely. +static __aicore__ void residual_rms_cast_2( + __gm__ bfloat16_t *v1, __gm__ float *v2, __gm__ float *v3, __gm__ float *v4, __gm__ float *v5, int64_t v6, + int64_t v7 +) { + SaturationMode v8 = SaturationMode::OFF; + RoundMode v9 = RoundMode::CAST_ROUND; + const int64_t v10 = 256; + const int64_t v11 = 2; + const int64_t v12 = 4; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 0; + const int64_t v17 = 51200; + const int64_t v18 = 34816; + const int64_t v19 = 33792; + const int64_t v20 = 17408; + const int64_t v21 = 1024; + using T = float; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + for (size_t v22 = (size_t)v16; v22 < ((size_t)v12); v22 += (size_t)v11) { + int64_t v23 = (int64_t)((uint64_t)((int64_t)v22) * (uint64_t)v10); + int64_t v24 = (int64_t)((uint64_t)v6 + (uint64_t)v23); + int64_t v25 = (int64_t)((uint64_t)v6 + (uint64_t)((int64_t)(uint64_t)v23 + (uint64_t)v10)); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v26 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v27 = (uint64_t)v21; + TASSIGN(v26, v27); + pto::Shape<1, 1, 1, 16, 256> v28 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v29 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v30 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v24 * v13), v28, v29 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + TLOAD(v26, v30); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v31 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v32 = (uint64_t)v20; + TASSIGN(v31, v32); + pto::Shape<1, 1, 1, 16, 256> v33 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v34 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v35 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v24 * v13), v33, v34 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + TLOAD(v31, v35); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v36 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v37 = (uint64_t)v19; + TASSIGN(v36, v37); + pto::Shape<1, 1, 1, 1, 256> v38 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v39 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v40 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v24 * v13), v38, v39 + ); + TLOAD(v36, v40); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v41 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v42 = (uint64_t)v18; + TASSIGN(v41, v42); + pto::Shape<1, 1, 1, 16, 256> v43 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v44 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v45 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v25 * v13), v43, v44 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + TLOAD(v41, v45); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v46 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v47 = (uint64_t)v17; + TASSIGN(v46, v47); + pto::Shape<1, 1, 1, 16, 256> v48 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v49 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v50 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v25 * v13), v48, v49 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + TLOAD(v46, v50); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v52 = (uint64_t)v16; + TASSIGN(v51, v52); + pto::Shape<1, 1, 1, 1, 256> v53 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v54 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v55 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v25 * v13), v53, v54 + ); + TLOAD(v51, v55); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v56 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v57 = (uint64_t)v21; + TASSIGN(v56, v57); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + TADD(v56, v26, v31); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v59 = (uint64_t)v20; + TASSIGN(v58, v59); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TCOLEXPANDMUL(v58, v56, v36); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v60 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v61 = (uint64_t)v20; + TASSIGN(v60, v61); + pipe_barrier(PIPE_V); + TCVT(v60, v58, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + pto::Shape<1, 1, 1, 16, 256> v62 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v63 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v64 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + v23 * v13), v62, v63 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + pipe_barrier(PIPE_MTE3); + TSTORE(v64, v56); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + pto::Shape<1, 1, 1, 16, 256> v65 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v66 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + v23 * v13), v65, v66 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(v67, v60); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v69 = (uint64_t)v18; + TASSIGN(v68, v69); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + TADD(v68, v41, v46); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v70 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v71 = (uint64_t)v17; + TASSIGN(v70, v71); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + TCOLEXPANDMUL(v70, v68, v51); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v72 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v73 = (uint64_t)v17; + TASSIGN(v72, v73); + pipe_barrier(PIPE_V); + TCVT(v72, v70, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + pto::Shape<1, 1, 1, 16, 256> v74 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v75 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v76 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + (v23 + v10) * v13), v74, v75 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + pipe_barrier(PIPE_MTE3); + TSTORE(v76, v68); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + pto::Shape<1, 1, 1, 16, 256> v77 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v78 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v79 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + (v23 + v10) * v13), v77, v78 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + TSTORE(v79, v72); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + } + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); +#endif // __DAV_VEC__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_norm_in_inline71__rv_v8 + __gm__ Tensor *mlp_norm_in_inline71__rv_v8_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_norm_in_inline71__rv_v8 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_norm_in_inline71__rv_v8_tensor->buffer.addr) + + mlp_norm_in_inline71__rv_v8_tensor->start_offset; + + // Unpack tensor: post_norm_partial_inline118__rv_v8 + __gm__ Tensor *post_norm_partial_inline118__rv_v8_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *post_norm_partial_inline118__rv_v8 = + reinterpret_cast<__gm__ float *>(post_norm_partial_inline118__rv_v8_tensor->buffer.addr) + + post_norm_partial_inline118__rv_v8_tensor->start_offset; + + // Unpack tensor: attn_proj_fp32_inline220__ssa_v7 + __gm__ Tensor *attn_proj_fp32_inline220__ssa_v7_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *attn_proj_fp32_inline220__ssa_v7 = + reinterpret_cast<__gm__ float *>(attn_proj_fp32_inline220__ssa_v7_tensor->buffer.addr) + + attn_proj_fp32_inline220__ssa_v7_tensor->start_offset; + + // Unpack tensor: cur__iter_v6 + __gm__ Tensor *cur__iter_v6_tensor = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ float *cur__iter_v6 = + reinterpret_cast<__gm__ float *>(cur__iter_v6_tensor->buffer.addr) + cur__iter_v6_tensor->start_offset; + + // Unpack tensor: post_rms_weight__ssa_v0 + __gm__ Tensor *post_rms_weight__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ float *post_rms_weight__ssa_v0 = + reinterpret_cast<__gm__ float *>(post_rms_weight__ssa_v0_tensor->buffer.addr) + + post_rms_weight__ssa_v0_tensor->start_offset; + + // Unpack scalar: k_base_inline111__ssa_v3 + union { + uint64_t u64; + int64_t val; + } k_base_inline111__ssa_v3_conv; + k_base_inline111__ssa_v3_conv.u64 = args[5]; + int64_t k_base_inline111__ssa_v3 = k_base_inline111__ssa_v3_conv.val; + + // Unpack scalar: i__idx_v0 + union { + uint64_t u64; + int64_t val; + } i__idx_v0_conv; + i__idx_v0_conv.u64 = args[6]; + int64_t i__idx_v0 = i__idx_v0_conv.val; + + // Forward to ptoas-generated function + residual_rms_cast_2( + mlp_norm_in_inline71__rv_v8, post_norm_partial_inline118__rv_v8, attn_proj_fp32_inline220__ssa_v7, cur__iter_v6, + post_rms_weight__ssa_v0, k_base_inline111__ssa_v3, i__idx_v0 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_3.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_3.cpp new file mode 100644 index 0000000000..b11e23b73b --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/residual_rms_cast_3.cpp @@ -0,0 +1,369 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: residual_rms_cast_3 + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +// The two stores below address the band relatively: the task's `mlp_norm_in` and +// `post_norm_partial` arguments are views narrowed to [k_base, k_base + 1024), so +// their base already carries k_base and only the in-band offset is added. The +// loads keep the parent views and so still index absolutely. +static __aicore__ void residual_rms_cast_3( + __gm__ bfloat16_t *v1, __gm__ float *v2, __gm__ float *v3, __gm__ float *v4, __gm__ float *v5, int64_t v6, + int64_t v7 +) { + SaturationMode v8 = SaturationMode::OFF; + RoundMode v9 = RoundMode::CAST_ROUND; + const int64_t v10 = 256; + const int64_t v11 = 2; + const int64_t v12 = 4; + const int64_t v13 = 1; + const int64_t v14 = 5120; + const int64_t v15 = 16; + const int64_t v16 = 0; + const int64_t v17 = 51200; + const int64_t v18 = 34816; + const int64_t v19 = 33792; + const int64_t v20 = 17408; + const int64_t v21 = 1024; + using T = float; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + for (size_t v22 = (size_t)v16; v22 < ((size_t)v12); v22 += (size_t)v11) { + int64_t v23 = (int64_t)((uint64_t)((int64_t)v22) * (uint64_t)v10); + int64_t v24 = (int64_t)((uint64_t)v6 + (uint64_t)v23); + int64_t v25 = (int64_t)((uint64_t)v6 + (uint64_t)((int64_t)(uint64_t)v23 + (uint64_t)v10)); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v26 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v27 = (uint64_t)v21; + TASSIGN(v26, v27); + pto::Shape<1, 1, 1, 16, 256> v28 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v29 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v30 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v24 * v13), v28, v29 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + TLOAD(v26, v30); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v31 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v32 = (uint64_t)v20; + TASSIGN(v31, v32); + pto::Shape<1, 1, 1, 16, 256> v33 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v34 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v35 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v24 * v13), v33, v34 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + TLOAD(v31, v35); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v36 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v37 = (uint64_t)v19; + TASSIGN(v36, v37); + pto::Shape<1, 1, 1, 1, 256> v38 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v39 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v40 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v24 * v13), v38, v39 + ); + TLOAD(v36, v40); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v41 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v42 = (uint64_t)v18; + TASSIGN(v41, v42); + pto::Shape<1, 1, 1, 16, 256> v43 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v44 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v45 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v3 + (v16 + v16 * v14 + v25 * v13), v43, v44 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + TLOAD(v41, v45); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v46 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v47 = (uint64_t)v17; + TASSIGN(v46, v47); + pto::Shape<1, 1, 1, 16, 256> v48 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v49 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v50 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v4 + (v16 + v16 * v14 + v25 * v13), v48, v49 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + TLOAD(v46, v50); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v51 = Tile< + TileType::Vec, float, 1, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v13, v10); + uint64_t v52 = (uint64_t)v16; + TASSIGN(v51, v52); + pto::Shape<1, 1, 1, 1, 256> v53 = pto::Shape<1, 1, 1, 1, 256>(); + pto::Stride<5120, 5120, 5120, 5120, 1> v54 = pto::Stride<5120, 5120, 5120, 5120, 1>(); + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND> v55 = + GlobalTensor, pto::Stride<5120, 5120, 5120, 5120, 1>, pto::Layout::ND>( + v5 + (v16 + v7 * v14 + v25 * v13), v53, v54 + ); + TLOAD(v51, v55); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v56 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v57 = (uint64_t)v21; + TASSIGN(v56, v57); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + TADD(v56, v26, v31); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v58 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v59 = (uint64_t)v20; + TASSIGN(v58, v59); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TCOLEXPANDMUL(v58, v56, v36); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v60 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v61 = (uint64_t)v20; + TASSIGN(v60, v61); + pipe_barrier(PIPE_V); + TCVT(v60, v58, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + pto::Shape<1, 1, 1, 16, 256> v62 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v63 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v64 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + v23 * v13), v62, v63 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + pipe_barrier(PIPE_MTE3); + TSTORE(v64, v56); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + pto::Shape<1, 1, 1, 16, 256> v65 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v66 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v67 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + v23 * v13), v65, v66 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + TSTORE(v67, v60); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v68 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v69 = (uint64_t)v18; + TASSIGN(v68, v69); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + TADD(v68, v41, v46); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v70 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v71 = (uint64_t)v17; + TASSIGN(v70, v71); + pipe_barrier(PIPE_V); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + TCOLEXPANDMUL(v70, v68, v51); + set_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v72 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v15, v10); + uint64_t v73 = (uint64_t)v17; + TASSIGN(v72, v73); + pipe_barrier(PIPE_V); + TCVT(v72, v70, v9, v8); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + pto::Shape<1, 1, 1, 16, 256> v74 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v75 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v76 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v2 + (v16 + v16 * v14 + (v23 + v10) * v13), v74, v75 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID2); + pipe_barrier(PIPE_MTE3); + TSTORE(v76, v68); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + pto::Shape<1, 1, 1, 16, 256> v77 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<81920, 81920, 81920, 5120, 1> v78 = pto::Stride<81920, 81920, 81920, 5120, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND> + v79 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<81920, 81920, 81920, 5120, 1>, pto::Layout::ND>( + v1 + (v16 + v16 * v14 + (v23 + v10) * v13), v77, v78 + ); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID3); + TSTORE(v79, v72); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); + } + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID2); + wait_flag(PIPE_V, PIPE_MTE2, EVENT_ID1); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID3); +#endif // __DAV_VEC__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: mlp_norm_in_inline71__rv_v11 + __gm__ Tensor *mlp_norm_in_inline71__rv_v11_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ bfloat16_t *mlp_norm_in_inline71__rv_v11 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_norm_in_inline71__rv_v11_tensor->buffer.addr) + + mlp_norm_in_inline71__rv_v11_tensor->start_offset; + + // Unpack tensor: post_norm_partial_inline118__rv_v11 + __gm__ Tensor *post_norm_partial_inline118__rv_v11_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ float *post_norm_partial_inline118__rv_v11 = + reinterpret_cast<__gm__ float *>(post_norm_partial_inline118__rv_v11_tensor->buffer.addr) + + post_norm_partial_inline118__rv_v11_tensor->start_offset; + + // Unpack tensor: attn_proj_fp32_inline220__ssa_v7 + __gm__ Tensor *attn_proj_fp32_inline220__ssa_v7_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *attn_proj_fp32_inline220__ssa_v7 = + reinterpret_cast<__gm__ float *>(attn_proj_fp32_inline220__ssa_v7_tensor->buffer.addr) + + attn_proj_fp32_inline220__ssa_v7_tensor->start_offset; + + // Unpack tensor: cur__iter_v6 + __gm__ Tensor *cur__iter_v6_tensor = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ float *cur__iter_v6 = + reinterpret_cast<__gm__ float *>(cur__iter_v6_tensor->buffer.addr) + cur__iter_v6_tensor->start_offset; + + // Unpack tensor: post_rms_weight__ssa_v0 + __gm__ Tensor *post_rms_weight__ssa_v0_tensor = reinterpret_cast<__gm__ Tensor *>(args[4]); + __gm__ float *post_rms_weight__ssa_v0 = + reinterpret_cast<__gm__ float *>(post_rms_weight__ssa_v0_tensor->buffer.addr) + + post_rms_weight__ssa_v0_tensor->start_offset; + + // Unpack scalar: k_base_inline111__ssa_v4 + union { + uint64_t u64; + int64_t val; + } k_base_inline111__ssa_v4_conv; + k_base_inline111__ssa_v4_conv.u64 = args[5]; + int64_t k_base_inline111__ssa_v4 = k_base_inline111__ssa_v4_conv.val; + + // Unpack scalar: i__idx_v0 + union { + uint64_t u64; + int64_t val; + } i__idx_v0_conv; + i__idx_v0_conv.u64 = args[6]; + int64_t i__idx_v0 = i__idx_v0_conv.val; + + // Forward to ptoas-generated function + residual_rms_cast_3( + mlp_norm_in_inline71__rv_v11, post_norm_partial_inline118__rv_v11, attn_proj_fp32_inline220__ssa_v7, + cur__iter_v6, post_rms_weight__ssa_v0, k_base_inline111__ssa_v4, i__idx_v0 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/silu.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/silu.cpp new file mode 100644 index 0000000000..edb1f7926f --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/aiv/silu.cpp @@ -0,0 +1,422 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Kernel Function: silu + +#include + +#ifndef __gm__ +#define __gm__ +#endif + +#ifndef __aicore__ +#if defined(__CPU_SIM) +#define __aicore__ +#else +#define __aicore__ [aicore] +#endif +#endif + +#include +#include "tensor.h" + +using namespace pto; + +// --- ptoas-generated code --- + +enum class PTOAutoSyncTailMode : int { + kBarrierAll = 0, + kSetWaitMte3ToSEvent0 = 1, +}; + +static __aicore__ inline void ptoas_auto_sync_tail(PTOAutoSyncTailMode mode = PTOAutoSyncTailMode::kBarrierAll) { + switch (mode) { + case PTOAutoSyncTailMode::kSetWaitMte3ToSEvent0: + set_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID0); + break; + case PTOAutoSyncTailMode::kBarrierAll: + default: + pipe_barrier(PIPE_ALL); + break; + } +} + +static __aicore__ void silu(__gm__ float *v1, __gm__ bfloat16_t *v2, __gm__ float *v3, __gm__ float *v4, int64_t v5) { + SaturationMode v6 = SaturationMode::OFF; + RoundMode v7 = RoundMode::CAST_ROUND; + const float v8 = 1.0f; + const int64_t v9 = 256; + const int64_t v10 = 2; + const int64_t v11 = 4; + const int64_t v12 = 17408; + const int64_t v13 = 1; + const int64_t v14 = 16; + const int64_t v15 = 49152; + const int64_t v16 = 32768; + const int64_t v17 = 16384; + const int64_t v18 = 0; + const int64_t v19 = 114752; + const int64_t v20 = 98368; + const int64_t v21 = 81984; + const int64_t v22 = 65600; + const int64_t v23 = 65536; + using T = float; + +#if defined(__DAV_VEC__) + set_mask_norm(); + set_vector_mask(-1, -1); + Tile< + TileType::Vec, float, 16, 1, BLayout::ColMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v24 = Tile< + TileType::Vec, float, 16, 1, BLayout::ColMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v13); + uint64_t v25 = (uint64_t)v23; + TASSIGN(v24, v25); + pto::Shape<1, 1, 1, 16, 1> v26 = pto::Shape<1, 1, 1, 16, 1>(); + pto::Stride<16, 16, 16, 1, 16> v27 = pto::Stride<16, 16, 16, 1, 16>(); + GlobalTensor, pto::Stride<16, 16, 16, 1, 16>, pto::Layout::DN> v28 = + GlobalTensor, pto::Stride<16, 16, 16, 1, 16>, pto::Layout::DN>( + v1 + (v18 + v18 * v13 + v18 * v14), v26, v27 + ); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + TLOAD(v24, v28); + for (size_t v29 = (size_t)v18; v29 < ((size_t)v11); v29 += (size_t)v10) { + int64_t v30 = (int64_t)((uint64_t)((int64_t)v29) * (uint64_t)v9); + int64_t v31 = (int64_t)((uint64_t)v5 + (uint64_t)v30); + int64_t v32 = (int64_t)((uint64_t)v5 + (uint64_t)((int64_t)(uint64_t)v30 + (uint64_t)v9)); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v33 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v34 = (uint64_t)v22; + TASSIGN(v33, v34); + pto::Shape<1, 1, 1, 16, 256> v35 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v36 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v37 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v3 + (v18 + v18 * v12 + (v31 - v5) * v13), v35, v36 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + TLOAD(v33, v37); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v38 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v39 = (uint64_t)v21; + TASSIGN(v38, v39); + pto::Shape<1, 1, 1, 16, 256> v40 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v41 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v42 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v4 + (v18 + v18 * v12 + (v31 - v5) * v13), v40, v41 + ); + TLOAD(v38, v42); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v43 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v44 = (uint64_t)v20; + TASSIGN(v43, v44); + pto::Shape<1, 1, 1, 16, 256> v45 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v46 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v47 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v3 + (v18 + v18 * v12 + (v32 - v5) * v13), v45, v46 + ); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + TLOAD(v43, v47); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v48 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v49 = (uint64_t)v19; + TASSIGN(v48, v49); + pto::Shape<1, 1, 1, 16, 256> v50 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v51 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v52 = GlobalTensor< + float, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND>( + v4 + (v18 + v18 * v12 + (v32 - v5) * v13), v50, v51 + ); + TLOAD(v48, v52); + set_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v53 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v54 = (uint64_t)v22; + TASSIGN(v53, v54); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0); + TROWEXPANDMUL(v53, v33, v24); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v55 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v56 = (uint64_t)v21; + TASSIGN(v55, v56); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID1); + TROWEXPANDMUL(v55, v38, v24); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v57 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v58 = (uint64_t)v18; + TASSIGN(v57, v58); + pipe_barrier(PIPE_V); + TNEG(v57, v53); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v59 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v60 = (uint64_t)v18; + TASSIGN(v59, v60); + pipe_barrier(PIPE_V); + TEXP(v59, v57); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v61 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v62 = (uint64_t)v18; + TASSIGN(v61, v62); + pipe_barrier(PIPE_V); + TADDS(v61, v59, v8); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v63 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v64 = (uint64_t)v17; + TASSIGN(v63, v64); + pipe_barrier(PIPE_V); + TRECIP(v63, v61); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v65 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v66 = (uint64_t)v22; + TASSIGN(v65, v66); + pipe_barrier(PIPE_V); + TMUL(v65, v53, v63); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v67 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v68 = (uint64_t)v22; + TASSIGN(v67, v68); + pipe_barrier(PIPE_V); + TMUL(v67, v65, v55); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v69 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v70 = (uint64_t)v22; + TASSIGN(v69, v70); + pipe_barrier(PIPE_V); + TCVT(v69, v67, v7, v6); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + pto::Shape<1, 1, 1, 16, 256> v71 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v72 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v73 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, + pto::Layout::ND>(v2 + (v18 + v18 * v12 + (v31 - v5) * v13), v71, v72); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0); + pipe_barrier(PIPE_MTE3); + TSTORE(v73, v69); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v74 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v75 = (uint64_t)v20; + TASSIGN(v74, v75); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID2); + TROWEXPANDMUL(v74, v43, v24); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v76 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v77 = (uint64_t)v19; + TASSIGN(v76, v77); + wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID3); + TROWEXPANDMUL(v76, v48, v24); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v78 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v79 = (uint64_t)v16; + TASSIGN(v78, v79); + pipe_barrier(PIPE_V); + TNEG(v78, v74); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v80 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v81 = (uint64_t)v16; + TASSIGN(v80, v81); + pipe_barrier(PIPE_V); + TEXP(v80, v78); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v82 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v83 = (uint64_t)v16; + TASSIGN(v82, v83); + pipe_barrier(PIPE_V); + TADDS(v82, v80, v8); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v84 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v85 = (uint64_t)v15; + TASSIGN(v84, v85); + pipe_barrier(PIPE_V); + TRECIP(v84, v82); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v86 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v87 = (uint64_t)v20; + TASSIGN(v86, v87); + pipe_barrier(PIPE_V); + TMUL(v86, v74, v84); + Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v88 = Tile< + TileType::Vec, float, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v89 = (uint64_t)v20; + TASSIGN(v88, v89); + pipe_barrier(PIPE_V); + TMUL(v88, v86, v76); + Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null> + v90 = Tile< + TileType::Vec, bfloat16_t, 16, 256, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Null, + CompactMode::Null>(v14, v9); + uint64_t v91 = (uint64_t)v20; + TASSIGN(v90, v91); + pipe_barrier(PIPE_V); + TCVT(v90, v88, v7, v6); + set_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + pto::Shape<1, 1, 1, 16, 256> v92 = pto::Shape<1, 1, 1, 16, 256>(); + pto::Stride<278528, 278528, 278528, 17408, 1> v93 = pto::Stride<278528, 278528, 278528, 17408, 1>(); + GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, pto::Layout::ND> + v94 = GlobalTensor< + bfloat16_t, pto::Shape<1, 1, 1, 16, 256>, pto::Stride<278528, 278528, 278528, 17408, 1>, + pto::Layout::ND>(v2 + (v18 + v18 * v12 + (v32 - v5) * v13), v92, v93); + wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID1); + pipe_barrier(PIPE_MTE3); + TSTORE(v94, v90); + set_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); + } + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID0); + wait_flag(PIPE_MTE3, PIPE_MTE2, EVENT_ID1); +#endif // __DAV_VEC__ + + ptoas_auto_sync_tail(PTOAutoSyncTailMode::kBarrierAll); + return; +} +// --- Kernel entry point --- +extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) { + // Unpack tensor: inv_rms_tile_inline126__ssa_v1 + __gm__ Tensor *inv_rms_tile_inline126__ssa_v1_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]); + __gm__ float *inv_rms_tile_inline126__ssa_v1 = + reinterpret_cast<__gm__ float *>(inv_rms_tile_inline126__ssa_v1_tensor->buffer.addr) + + inv_rms_tile_inline126__ssa_v1_tensor->start_offset; + + // Unpack tensor: mlp_tile_inline149__iter_v1 + __gm__ Tensor *mlp_tile_inline149__iter_v1_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]); + __gm__ bfloat16_t *mlp_tile_inline149__iter_v1 = + reinterpret_cast<__gm__ bfloat16_t *>(mlp_tile_inline149__iter_v1_tensor->buffer.addr) + + mlp_tile_inline149__iter_v1_tensor->start_offset; + + // Unpack tensor: gate_acc_all_inline203__rv_v10 + __gm__ Tensor *gate_acc_all_inline203__rv_v10_tensor = reinterpret_cast<__gm__ Tensor *>(args[2]); + __gm__ float *gate_acc_all_inline203__rv_v10 = + reinterpret_cast<__gm__ float *>(gate_acc_all_inline203__rv_v10_tensor->buffer.addr) + + gate_acc_all_inline203__rv_v10_tensor->start_offset; + + // Unpack tensor: up_acc_all_inline303__rv_v10 + __gm__ Tensor *up_acc_all_inline303__rv_v10_tensor = reinterpret_cast<__gm__ Tensor *>(args[3]); + __gm__ float *up_acc_all_inline303__rv_v10 = + reinterpret_cast<__gm__ float *>(up_acc_all_inline303__rv_v10_tensor->buffer.addr) + + up_acc_all_inline303__rv_v10_tensor->start_offset; + + // Unpack scalar: n0_inline122__ssa_v7 + union { + uint64_t u64; + int64_t val; + } n0_inline122__ssa_v7_conv; + n0_inline122__ssa_v7_conv.u64 = args[4]; + int64_t n0_inline122__ssa_v7 = n0_inline122__ssa_v7_conv.val; + + // Forward to ptoas-generated function + silu( + inv_rms_tile_inline126__ssa_v1, mlp_tile_inline149__iter_v1, gate_acc_all_inline203__rv_v10, + up_acc_all_inline303__rv_v10, n0_inline122__ssa_v7 + ); +} diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/orchestration/decode_fwd_layers.cpp b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/orchestration/decode_fwd_layers.cpp new file mode 100644 index 0000000000..901e415b0b --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/kernels/orchestration/decode_fwd_layers.cpp @@ -0,0 +1,2265 @@ +/* + * Copyright (c) PyPTO Contributors. + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of + * CANN Open Software License Agreement Version 2.0 (the "License"). + * Please refer to the License for details. You may not use this file except in compliance with the License. + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. + * See LICENSE in the root of the software repository for the full text of the License. + * ----------------------------------------------------------------------------------------------------------- + */ +// Orchestration Function: decode_fwd_layers +// Generated by PyPTO IR Compiler + +#include "runtime.h" +#include + +#include +#include +#include + +#include "orchestration_api.h" + +extern "C" { + +__attribute__((visibility("default"))) OrchestrationConfig aicpu_orchestration_config(const ChipTaskArgs &orch_args) { + (void)orch_args; + return OrchestrationConfig{ + .expected_arg_count = 20, + }; +} + +__attribute__((visibility("default"))) void aicpu_orchestration_entry(const ChipTaskArgs &orch_args) { + // External tensors + const simpler::tmr::Tensor &ext_hidden_states = orch_args.tensor(0).ref(); + const simpler::tmr::Tensor &ext_input_rms_weight = orch_args.tensor(1).ref(); + const simpler::tmr::Tensor &ext_wq = orch_args.tensor(2).ref(); + const simpler::tmr::Tensor &ext_wk = orch_args.tensor(3).ref(); + const simpler::tmr::Tensor &ext_wv = orch_args.tensor(4).ref(); + const simpler::tmr::Tensor &ext_q_norm_weight = orch_args.tensor(5).ref(); + const simpler::tmr::Tensor &ext_k_norm_weight = orch_args.tensor(6).ref(); + const simpler::tmr::Tensor &ext_seq_lens = orch_args.tensor(7).ref(); + const simpler::tmr::Tensor &ext_block_table = orch_args.tensor(8).ref(); + const simpler::tmr::Tensor &ext_slot_mapping = orch_args.tensor(9).ref(); + const simpler::tmr::Tensor &ext_rope_cos = orch_args.tensor(10).ref(); + const simpler::tmr::Tensor &ext_rope_sin = orch_args.tensor(11).ref(); + const simpler::tmr::Tensor &ext_k_cache = orch_args.tensor(12).ref(); + const simpler::tmr::Tensor &ext_v_cache = orch_args.tensor(13).ref(); + const simpler::tmr::Tensor &ext_wo = orch_args.tensor(14).ref(); + const simpler::tmr::Tensor &ext_w_gate = orch_args.tensor(15).ref(); + const simpler::tmr::Tensor &ext_w_up = orch_args.tensor(16).ref(); + const simpler::tmr::Tensor &ext_w_down = orch_args.tensor(17).ref(); + const simpler::tmr::Tensor &ext_post_rms_weight = orch_args.tensor(18).ref(); + const simpler::tmr::Tensor &ext_out = orch_args.tensor(19).ref(); + + // Dynamic-dim symbols (extent of the declaring argument) + int64_t BLOCK_TABLE_FLAT_DYN = (int64_t)orch_args.tensor(8).ref().shapes[0]; + int64_t KV_CACHE_ROWS_DYN = (int64_t)orch_args.tensor(12).ref().shapes[0]; + + SIMPLER_SCOPE() { + uint32_t pa_metadata_ci_shapes[1] = {27840}; + TensorCreateInfo pa_metadata_ci(pa_metadata_ci_shapes, 1, DataType::UINT8); + uint32_t pa_workspace_ci_shapes[1] = {66132544}; + TensorCreateInfo pa_workspace_ci(pa_workspace_ci_shapes, 1, DataType::UINT8); + uint32_t cur_ci_shapes[2] = {16, 5120}; + TensorCreateInfo cur_ci(cur_ci_shapes, 2, DataType::FLOAT32); + uint32_t normed_ci_shapes[2] = {16, 5120}; + TensorCreateInfo normed_ci(normed_ci_shapes, 2, DataType::BFLOAT16); + TaskOutputTensors alloc_0 = alloc_tensors(pa_metadata_ci, pa_workspace_ci, cur_ci, normed_ci); + const simpler::tmr::Tensor &pa_metadata = alloc_0.get_ref(0); + const simpler::tmr::Tensor &pa_workspace = alloc_0.get_ref(1); + const simpler::tmr::Tensor &cur = alloc_0.get_ref(2); + const simpler::tmr::Tensor &normed = alloc_0.get_ref(3); + int64_t pa_num_layers = 40; + int64_t pa_num_pages = (KV_CACHE_ROWS_DYN / (pa_num_layers * 1024)); + int64_t pa_max_blocks = (BLOCK_TABLE_FLAT_DYN / 16); + int32_t pa_num_pages_i32 = static_cast(pa_num_pages); + int32_t pa_max_blocks_i32 = static_cast(pa_max_blocks); + + // Spmd pa_tiling: paged_attention_tiling_cce + CoreTaskArgs params_t0; + params_t0.add_input(ext_seq_lens); + params_t0.add_output(pa_metadata); + params_t0.add_scalar(pa_max_blocks_i32); + params_t0.add_scalar(pa_num_pages_i32); + params_t0.launch_spec.set_block_num(1); + params_t0.set_allow_early_resolve(true); + TaskOutputTensors task_0_outs = rt_submit_aiv_task(0, params_t0); + TaskId tiling_tid_inline0 = task_0_outs.task_id(); + TaskId pa_tiling_tid = tiling_tid_inline0; + TaskId prev_out_tid[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + prev_out_tid[__init_i] = TaskId::invalid(); + + // Phase-fence barrier 0: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_0; + TaskOutputTensors phase_fence_barrier_0_outs = rt_submit_dummy_task(params_phase_fence_barrier_0); + TaskId t = phase_fence_barrier_0_outs.task_id(); + prev_out_tid[0] = t; + for (int64_t cb0 = 0; cb0 < 16; cb0 += 16) { + SIMPLER_SCOPE() { + // Task 1: copy_hidden + CoreTaskArgs params_t1; + params_t1.add_output(cur); + params_t1.add_input(ext_hidden_states); + params_t1.add_scalar(cb0); + TaskOutputTensors task_1_outs = rt_submit_aiv_task(1, params_t1); + TaskId ch_tid = task_1_outs.task_id(); + prev_out_tid[0] = ch_tid; + } + } + TaskId prev_normed_tid[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + prev_normed_tid[__init_i] = TaskId::invalid(); + SIMPLER_SCOPE(ScopeMode::AUTO) { + TaskId _submit_deps_buf[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf[__init_i] = TaskId::invalid(); + TaskId t__tmp_v3 = prev_out_tid[0]; + _submit_deps_buf[0] = t__tmp_v3; + + // Spmd x_gamma0_spmd: x_gamma0 + CoreTaskArgs params_t2; + params_t2.add_output(normed); + params_t2.add_input(cur); + params_t2.add_input(ext_input_rms_weight); + params_t2.launch_spec.set_block_num(5); + params_t2.set_allow_early_resolve(true); + TaskId params_t2_deps[1]; + uint32_t params_t2_deps_count = 0; + if (_submit_deps_buf[0].is_valid()) params_t2_deps[params_t2_deps_count++] = _submit_deps_buf[0]; + params_t2.set_dependencies(params_t2_deps, params_t2_deps_count); + TaskOutputTensors task_2_outs = rt_submit_aiv_task(2, params_t2); + TaskId xgamma_tid = task_2_outs.task_id(); + prev_normed_tid[0] = xgamma_tid; + } + simpler::tmr::Tensor cur__rv_v7 = cur; + simpler::tmr::Tensor normed__rv_v5 = normed; + for (int64_t i = 0; i < 40; i += 1) { + SIMPLER_SCOPE() { + uint32_t next_hidden_ci_shapes[2] = {16, 5120}; + TensorCreateInfo next_hidden_ci(next_hidden_ci_shapes, 2, DataType::FLOAT32); + uint32_t next_normed_ci_shapes[2] = {16, 5120}; + TensorCreateInfo next_normed_ci(next_normed_ci_shapes, 2, DataType::BFLOAT16); + uint32_t inv_rms_states_inline176_ci_shapes[2] = {16, 1}; + TensorCreateInfo inv_rms_states_inline176_ci(inv_rms_states_inline176_ci_shapes, 2, DataType::FLOAT32); + uint32_t q_proj_inline139_ci_shapes[2] = {16, 5120}; + TensorCreateInfo q_proj_inline139_ci(q_proj_inline139_ci_shapes, 2, DataType::FLOAT32); + uint32_t k_proj_inline135_ci_shapes[2] = {16, 1024}; + TensorCreateInfo k_proj_inline135_ci(k_proj_inline135_ci_shapes, 2, DataType::FLOAT32); + uint32_t v_proj_inline255_ci_shapes[2] = {16, 1024}; + TensorCreateInfo v_proj_inline255_ci(v_proj_inline255_ci_shapes, 2, DataType::FLOAT32); + uint32_t q_tnd_flat_inline127_ci_shapes[2] = {640, 128}; + TensorCreateInfo q_tnd_flat_inline127_ci(q_tnd_flat_inline127_ci_shapes, 2, DataType::BFLOAT16); + uint32_t attn_out_inline282_ci_shapes[2] = {16, 5120}; + TensorCreateInfo attn_out_inline282_ci(attn_out_inline282_ci_shapes, 2, DataType::BFLOAT16); + TaskOutputTensors alloc_1 = alloc_tensors( + next_hidden_ci, next_normed_ci, inv_rms_states_inline176_ci, q_proj_inline139_ci, + k_proj_inline135_ci, v_proj_inline255_ci, q_tnd_flat_inline127_ci, attn_out_inline282_ci + ); + const simpler::tmr::Tensor &next_hidden = alloc_1.get_ref(0); + const simpler::tmr::Tensor &next_normed = alloc_1.get_ref(1); + const simpler::tmr::Tensor &inv_rms_states_inline176 = alloc_1.get_ref(2); + const simpler::tmr::Tensor &q_proj_inline139 = alloc_1.get_ref(3); + const simpler::tmr::Tensor &k_proj_inline135 = alloc_1.get_ref(4); + const simpler::tmr::Tensor &v_proj_inline255 = alloc_1.get_ref(5); + const simpler::tmr::Tensor &q_tnd_flat_inline127 = alloc_1.get_ref(6); + const simpler::tmr::Tensor &attn_out_inline282 = alloc_1.get_ref(7); + int64_t next_gamma_idx = std::min((i + 1), 39); + int64_t layer_hidden_base_inline151 = (static_cast(i) * 5120); + int64_t layer_inter_base_inline107 = (static_cast(i) * 17408); + int64_t num_layers_actual_inline152 = 40; + int64_t t__tmp_v6 = (int64_t)orch_args.tensor(12).ref().shapes[0]; + int64_t layer_cache_rows_inline128 = (t__tmp_v6 / num_layers_actual_inline152); + int64_t layer_cache_base_inline193 = (static_cast(i) * layer_cache_rows_inline128); + uint32_t q_norm_w_inline124_offsets[2] = {static_cast(i), 0}; + uint32_t q_norm_w_inline124_shapes[2] = { + (q_norm_w_inline124_offsets[0] >= ext_q_norm_weight.shapes[0] ? + 0u : + std::min(1, ext_q_norm_weight.shapes[0] - q_norm_w_inline124_offsets[0])), + (q_norm_w_inline124_offsets[1] >= ext_q_norm_weight.shapes[1] ? + 0u : + std::min(128, ext_q_norm_weight.shapes[1] - q_norm_w_inline124_offsets[1])) + }; + simpler::tmr::Tensor q_norm_w_inline124 = + ext_q_norm_weight.view(q_norm_w_inline124_shapes, q_norm_w_inline124_offsets); + uint32_t k_norm_w_inline114_offsets[2] = {static_cast(i), 0}; + uint32_t k_norm_w_inline114_shapes[2] = { + (k_norm_w_inline114_offsets[0] >= ext_k_norm_weight.shapes[0] ? + 0u : + std::min(1, ext_k_norm_weight.shapes[0] - k_norm_w_inline114_offsets[0])), + (k_norm_w_inline114_offsets[1] >= ext_k_norm_weight.shapes[1] ? + 0u : + std::min(128, ext_k_norm_weight.shapes[1] - k_norm_w_inline114_offsets[1])) + }; + simpler::tmr::Tensor k_norm_w_inline114 = + ext_k_norm_weight.view(k_norm_w_inline114_shapes, k_norm_w_inline114_offsets); + TaskId down_tids_inline156[85]; + for (int64_t __init_i = 0; __init_i < 85; ++__init_i) + down_tids_inline156[__init_i] = TaskId::invalid(); + uint32_t down_acc_all_inline168_ci_shapes[2] = {16, 5120}; + TensorCreateInfo down_acc_all_inline168_ci(down_acc_all_inline168_ci_shapes, 2, DataType::FLOAT32); + uint32_t gate_acc_all_inline203_ci_shapes[2] = {16, 17408}; + TensorCreateInfo gate_acc_all_inline203_ci(gate_acc_all_inline203_ci_shapes, 2, DataType::FLOAT32); + uint32_t up_acc_all_inline303_ci_shapes[2] = {16, 17408}; + TensorCreateInfo up_acc_all_inline303_ci(up_acc_all_inline303_ci_shapes, 2, DataType::FLOAT32); + uint32_t attn_proj_fp32_inline220_ci_shapes[2] = {16, 5120}; + TensorCreateInfo attn_proj_fp32_inline220_ci(attn_proj_fp32_inline220_ci_shapes, 2, DataType::FLOAT32); + uint32_t post_norm_partial_inline118_ci_shapes[2] = {16, 5120}; + TensorCreateInfo post_norm_partial_inline118_ci( + post_norm_partial_inline118_ci_shapes, 2, DataType::FLOAT32 + ); + uint32_t mlp_norm_in_inline71_ci_shapes[2] = {16, 5120}; + TensorCreateInfo mlp_norm_in_inline71_ci(mlp_norm_in_inline71_ci_shapes, 2, DataType::BFLOAT16); + uint32_t inv_rms_tile_inline126_ci_shapes[2] = {16, 1}; + TensorCreateInfo inv_rms_tile_inline126_ci(inv_rms_tile_inline126_ci_shapes, 2, DataType::FLOAT32); + uint32_t mlp_tile_inline149_ci_shapes[2] = {16, 17408}; + TensorCreateInfo mlp_tile_inline149_ci(mlp_tile_inline149_ci_shapes, 2, DataType::BFLOAT16); + // Private split-K partials: N stacked copies of the accumulator's [16, W] shape, + // so partial j is rows [j*16, j*16+16). Row stride stays W, which makes a partial's + // band view geometrically identical to the accumulator band the kernel used to + // atomically add into — the store keeps its Shape/Stride constants and only drops + // the atomic. + uint32_t down_part_all_ci_shapes[2] = {272, 5120}; + TensorCreateInfo down_part_all_ci(down_part_all_ci_shapes, 2, DataType::FLOAT32); + uint32_t out_part_all_ci_shapes[2] = {80, 5120}; + TensorCreateInfo out_part_all_ci(out_part_all_ci_shapes, 2, DataType::FLOAT32); + uint32_t gate_part_all_ci_shapes[2] = {80, 17408}; + TensorCreateInfo gate_part_all_ci(gate_part_all_ci_shapes, 2, DataType::FLOAT32); + uint32_t up_part_all_ci_shapes[2] = {80, 17408}; + TensorCreateInfo up_part_all_ci(up_part_all_ci_shapes, 2, DataType::FLOAT32); + TaskOutputTensors alloc_2 = alloc_tensors( + down_acc_all_inline168_ci, gate_acc_all_inline203_ci, up_acc_all_inline303_ci, + attn_proj_fp32_inline220_ci, post_norm_partial_inline118_ci, mlp_norm_in_inline71_ci, + inv_rms_tile_inline126_ci, mlp_tile_inline149_ci, down_part_all_ci, out_part_all_ci, + gate_part_all_ci, up_part_all_ci + ); + const simpler::tmr::Tensor &down_acc_all_inline168 = alloc_2.get_ref(0); + const simpler::tmr::Tensor &gate_acc_all_inline203 = alloc_2.get_ref(1); + const simpler::tmr::Tensor &up_acc_all_inline303 = alloc_2.get_ref(2); + const simpler::tmr::Tensor &attn_proj_fp32_inline220 = alloc_2.get_ref(3); + const simpler::tmr::Tensor &post_norm_partial_inline118 = alloc_2.get_ref(4); + const simpler::tmr::Tensor &mlp_norm_in_inline71 = alloc_2.get_ref(5); + const simpler::tmr::Tensor &inv_rms_tile_inline126 = alloc_2.get_ref(6); + const simpler::tmr::Tensor &mlp_tile_inline149 = alloc_2.get_ref(7); + const simpler::tmr::Tensor &down_part_all = alloc_2.get_ref(8); + const simpler::tmr::Tensor &out_part_all = alloc_2.get_ref(9); + const simpler::tmr::Tensor &gate_part_all = alloc_2.get_ref(10); + const simpler::tmr::Tensor &up_part_all = alloc_2.get_ref(11); + SIMPLER_SCOPE(ScopeMode::AUTO) { + // Phase-fence barrier 1: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_1; + TaskOutputTensors phase_fence_barrier_1_outs = rt_submit_dummy_task(params_phase_fence_barrier_1); + TaskId seed_dummy_inline49 = phase_fence_barrier_1_outs.task_id(); + TaskId prev_normed_seed_deps_inline120[2]; + for (int64_t __init_i = 0; __init_i < 2; ++__init_i) + prev_normed_seed_deps_inline120[__init_i] = TaskId::invalid(); + TaskId t__tmp_v7 = prev_normed_tid[0]; + prev_normed_seed_deps_inline120[0] = t__tmp_v7; + prev_normed_seed_deps_inline120[1] = seed_dummy_inline49; + + // Task 3: attn_out_seed + CoreTaskArgs params_t3; + params_t3.add_input(attn_out_inline282); + params_t3.set_allow_early_resolve(true); + TaskOutputTensors task_3_outs = rt_submit_aiv_task(3, params_t3); + TaskId attn_out_seed_tid_inline116 = task_3_outs.task_id(); + TaskId _submit_deps_buf_inline42[2]; + for (int64_t __init_i = 0; __init_i < 2; ++__init_i) + _submit_deps_buf_inline42[__init_i] = TaskId::invalid(); + TaskId t__tmp_v8 = prev_normed_seed_deps_inline120[0]; + _submit_deps_buf_inline42[0] = t__tmp_v8; + TaskId t__tmp_v9 = prev_normed_seed_deps_inline120[1]; + _submit_deps_buf_inline42[1] = t__tmp_v9; + + // Task 4: rms_recip + CoreTaskArgs params_t4; + params_t4.add_input(cur__rv_v7); + params_t4.add_inout(inv_rms_states_inline176); + TaskId params_t4_deps[2]; + uint32_t params_t4_deps_count = 0; + if (_submit_deps_buf_inline42[0].is_valid()) + params_t4_deps[params_t4_deps_count++] = _submit_deps_buf_inline42[0]; + if (_submit_deps_buf_inline42[1].is_valid()) + params_t4_deps[params_t4_deps_count++] = _submit_deps_buf_inline42[1]; + params_t4.set_dependencies(params_t4_deps, params_t4_deps_count); + params_t4.set_allow_early_resolve(true); + TaskOutputTensors task_4_outs = rt_submit_aiv_task(4, params_t4); + TaskId rms_tid_inline148 = task_4_outs.task_id(); + + // Task 5: q_seed + CoreTaskArgs params_t5; + params_t5.add_inout(q_proj_inline139); + params_t5.set_allow_early_resolve(true); + TaskOutputTensors task_5_outs = rt_submit_aiv_task(5, params_t5); + TaskId q_seed_tid_inline162 = task_5_outs.task_id(); + TaskId prev_normed_q_deps_inline105[2]; + for (int64_t __init_i = 0; __init_i < 2; ++__init_i) + prev_normed_q_deps_inline105[__init_i] = TaskId::invalid(); + TaskId t__tmp_v17 = prev_normed_tid[0]; + prev_normed_q_deps_inline105[0] = t__tmp_v17; + prev_normed_q_deps_inline105[1] = q_seed_tid_inline162; + TaskId _submit_deps_buf_inline182[2]; + for (int64_t __init_i = 0; __init_i < 2; ++__init_i) + _submit_deps_buf_inline182[__init_i] = TaskId::invalid(); + TaskId t__tmp_v18 = prev_normed_q_deps_inline105[0]; + _submit_deps_buf_inline182[0] = t__tmp_v18; + TaskId t__tmp_v19 = prev_normed_q_deps_inline105[1]; + _submit_deps_buf_inline182[1] = t__tmp_v19; + + // Spmd q_proj_spmd: q_proj + CoreTaskArgs params_t6; + params_t6.add_inout(q_proj_inline139); + params_t6.add_input(normed__rv_v5); + params_t6.add_input(ext_wq); + params_t6.add_scalar(layer_hidden_base_inline151); + params_t6.launch_spec.set_block_num(50); + params_t6.set_allow_early_resolve(true); + TaskId params_t6_deps[2]; + uint32_t params_t6_deps_count = 0; + if (_submit_deps_buf_inline182[0].is_valid()) + params_t6_deps[params_t6_deps_count++] = _submit_deps_buf_inline182[0]; + if (_submit_deps_buf_inline182[1].is_valid()) + params_t6_deps[params_t6_deps_count++] = _submit_deps_buf_inline182[1]; + params_t6.set_dependencies(params_t6_deps, params_t6_deps_count); + TaskOutputTensors task_6_outs = rt_submit_aic_task(6, params_t6); + TaskId q_proj_tid_inline183 = task_6_outs.task_id(); + TaskId _submit_deps_buf_inline261[2]; + for (int64_t __init_i = 0; __init_i < 2; ++__init_i) + _submit_deps_buf_inline261[__init_i] = TaskId::invalid(); + TaskId t__tmp_v24 = prev_normed_seed_deps_inline120[0]; + _submit_deps_buf_inline261[0] = t__tmp_v24; + TaskId t__tmp_v25 = prev_normed_seed_deps_inline120[1]; + _submit_deps_buf_inline261[1] = t__tmp_v25; + + // Task 7: kv_seed + CoreTaskArgs params_t7; + params_t7.add_inout(k_proj_inline135); + params_t7.add_inout(v_proj_inline255); + TaskId params_t7_deps[2]; + uint32_t params_t7_deps_count = 0; + if (_submit_deps_buf_inline261[0].is_valid()) + params_t7_deps[params_t7_deps_count++] = _submit_deps_buf_inline261[0]; + if (_submit_deps_buf_inline261[1].is_valid()) + params_t7_deps[params_t7_deps_count++] = _submit_deps_buf_inline261[1]; + params_t7.set_dependencies(params_t7_deps, params_t7_deps_count); + TaskOutputTensors task_7_outs = rt_submit_aiv_task(7, params_t7); + TaskId kv_seed_tid_inline238 = task_7_outs.task_id(); + TaskId _submit_deps_buf_inline267[2]; + for (int64_t __init_i = 0; __init_i < 2; ++__init_i) + _submit_deps_buf_inline267[__init_i] = TaskId::invalid(); + TaskId t__tmp_v28 = prev_normed_seed_deps_inline120[0]; + _submit_deps_buf_inline267[0] = t__tmp_v28; + TaskId t__tmp_v29 = prev_normed_seed_deps_inline120[1]; + _submit_deps_buf_inline267[1] = t__tmp_v29; + + // Task 8: mlp_out_seed + CoreTaskArgs params_t8; + params_t8.add_inout(down_acc_all_inline168); + params_t8.add_inout(gate_acc_all_inline203); + params_t8.add_inout(up_acc_all_inline303); + params_t8.add_inout(attn_proj_fp32_inline220); + TaskId params_t8_deps[2]; + uint32_t params_t8_deps_count = 0; + if (_submit_deps_buf_inline267[0].is_valid()) + params_t8_deps[params_t8_deps_count++] = _submit_deps_buf_inline267[0]; + if (_submit_deps_buf_inline267[1].is_valid()) + params_t8_deps[params_t8_deps_count++] = _submit_deps_buf_inline267[1]; + params_t8.set_dependencies(params_t8_deps, params_t8_deps_count); + params_t8.set_allow_early_resolve(true); + TaskOutputTensors task_8_outs = rt_submit_aiv_task(8, params_t8); + TaskId mlp_out_seed_tid_inline206 = task_8_outs.task_id(); + + // Spmd k_proj_spmd: k_proj + CoreTaskArgs params_t9; + params_t9.add_inout(k_proj_inline135); + params_t9.add_input(normed__rv_v5); + params_t9.add_input(ext_wk); + params_t9.add_scalar(layer_hidden_base_inline151); + params_t9.launch_spec.set_block_num(10); + params_t9.set_allow_early_resolve(true); + TaskId params_t9_deps[1]; + uint32_t params_t9_deps_count = 0; + params_t9_deps[params_t9_deps_count++] = kv_seed_tid_inline238; + params_t9.set_dependencies(params_t9_deps, params_t9_deps_count); + TaskOutputTensors task_9_outs = rt_submit_aic_task(9, params_t9); + TaskId k_proj_tid_inline136 = task_9_outs.task_id(); + + // Spmd v_proj_spmd: v_proj + CoreTaskArgs params_t10; + params_t10.add_inout(v_proj_inline255); + params_t10.add_input(normed__rv_v5); + params_t10.add_input(ext_wv); + params_t10.add_scalar(layer_hidden_base_inline151); + params_t10.launch_spec.set_block_num(10); + params_t10.set_allow_early_resolve(true); + TaskId params_t10_deps[1]; + uint32_t params_t10_deps_count = 0; + params_t10_deps[params_t10_deps_count++] = kv_seed_tid_inline238; + params_t10.set_dependencies(params_t10_deps, params_t10_deps_count); + TaskOutputTensors task_10_outs = rt_submit_aic_task(10, params_t10); + TaskId v_proj_tid_inline63 = task_10_outs.task_id(); + uint32_t q_tnd_inline191_shapes[3] = {16, 40, 128}; + simpler::tmr::Tensor q_tnd_inline191 = q_tnd_flat_inline127.reshape(q_tnd_inline191_shapes, 3); + uint32_t attn_out_tnd_inline79_shapes[3] = {16, 40, 128}; + simpler::tmr::Tensor attn_out_tnd_inline79 = + attn_out_inline282.reshape(attn_out_tnd_inline79_shapes, 3); + int64_t attention_core_num_inline188 = 24; + + // Group paged_attention_rope_cce: MixedKernels (AIC + AIV lanes) + CoreTaskArgs params_t11; + params_t11.add_inout(attn_out_tnd_inline79); + params_t11.add_inout(q_tnd_inline191); + params_t11.add_inout(ext_k_cache); + params_t11.add_inout(ext_v_cache); + params_t11.add_input(ext_block_table); + params_t11.add_inout(pa_workspace); + params_t11.add_inout(pa_metadata); + params_t11.add_input(q_proj_inline139); + params_t11.add_input(k_proj_inline135); + params_t11.add_input(v_proj_inline255); + params_t11.add_input(q_norm_w_inline124); + params_t11.add_input(k_norm_w_inline114); + params_t11.add_input(ext_rope_cos); + params_t11.add_input(ext_rope_sin); + params_t11.add_input(inv_rms_states_inline176); + params_t11.add_input(ext_slot_mapping); + params_t11.add_input(ext_seq_lens); + params_t11.add_scalar(layer_cache_base_inline193); + MixedKernels mixed_11 = {11, 12, 12}; + params_t11.launch_spec.set_block_num(attention_core_num_inline188); + params_t11.launch_spec.set_require_sync_start(true); + params_t11.set_allow_early_resolve(true); + TaskId params_t11_deps[7]; + uint32_t params_t11_deps_count = 0; + params_t11_deps[params_t11_deps_count++] = q_proj_tid_inline183; + params_t11_deps[params_t11_deps_count++] = k_proj_tid_inline136; + params_t11_deps[params_t11_deps_count++] = v_proj_tid_inline63; + params_t11_deps[params_t11_deps_count++] = rms_tid_inline148; + params_t11_deps[params_t11_deps_count++] = tiling_tid_inline0; + params_t11_deps[params_t11_deps_count++] = attn_out_seed_tid_inline116; + params_t11_deps[params_t11_deps_count++] = mlp_out_seed_tid_inline206; + params_t11.set_dependencies(params_t11_deps, params_t11_deps_count); + TaskOutputTensors task_11_outs = rt_submit_task(mixed_11, params_t11); + const simpler::tmr::Tensor &attn_out_tnd_inline79__ssa_v1 = attn_out_tnd_inline79; + TaskId attn_done_tid_inline78 = task_11_outs.task_id(); + uint32_t attn_out_inline282__ssa_v4_shapes[2] = {16, 5120}; + simpler::tmr::Tensor attn_out_inline282__ssa_v4 = + attn_out_tnd_inline79__ssa_v1.reshape(attn_out_inline282__ssa_v4_shapes, 2); + TaskId silu_tids_inline265[17]; + for (int64_t __init_i = 0; __init_i < 17; ++__init_i) + silu_tids_inline265[__init_i] = TaskId::invalid(); + TaskId gate_tids_inline56[85]; + for (int64_t __init_i = 0; __init_i < 85; ++__init_i) + gate_tids_inline56[__init_i] = TaskId::invalid(); + TaskId up_tids_inline310[85]; + for (int64_t __init_i = 0; __init_i < 85; ++__init_i) + up_tids_inline310[__init_i] = TaskId::invalid(); + TaskId cast_tids_inline88[5]; + for (int64_t __init_i = 0; __init_i < 5; ++__init_i) + cast_tids_inline88[__init_i] = TaskId::invalid(); + TaskId gate_late_tids_inline249[5]; + for (int64_t __init_i = 0; __init_i < 5; ++__init_i) + gate_late_tids_inline249[__init_i] = TaskId::invalid(); + TaskId up_late_tids_inline69[5]; + for (int64_t __init_i = 0; __init_i < 5; ++__init_i) + up_late_tids_inline69[__init_i] = TaskId::invalid(); + TaskId out_tids_inline271[50]; + for (int64_t __init_i = 0; __init_i < 50; ++__init_i) + out_tids_inline271[__init_i] = TaskId::invalid(); + + // Phase-fence barrier 2: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_2; + TaskId params_phase_fence_barrier_2_deps[1]; + uint32_t params_phase_fence_barrier_2_deps_count = 0; + params_phase_fence_barrier_2_deps[params_phase_fence_barrier_2_deps_count++] = + attn_done_tid_inline78; + params_phase_fence_barrier_2.set_dependencies( + params_phase_fence_barrier_2_deps, params_phase_fence_barrier_2_deps_count + ); + TaskId out_proj_dummy_inline257 = TaskId::invalid(); + if (params_phase_fence_barrier_2_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_2_outs = + rt_submit_dummy_task(params_phase_fence_barrier_2); + out_proj_dummy_inline257 = phase_fence_barrier_2_outs.task_id(); + } + int64_t N_OUT_DIRECT_inline61 = 26; + for (int64_t out_idx_inline74 = 0; out_idx_inline74 < N_OUT_DIRECT_inline61; + out_idx_inline74 += 1) { + int64_t n_out_proj_inline185 = (out_idx_inline74 / 5); + int64_t k_split_out_inline66 = (out_idx_inline74 % 5); + int64_t n_op_inline64 = (n_out_proj_inline185 * 512); + int64_t k_op_inline266 = (k_split_out_inline66 * 1024); + + // Task 12: out_proj + CoreTaskArgs params_t12; + params_t12.add_input(attn_out_inline282__ssa_v4); + params_t12.add_input(ext_wo); + // Each split writes its own [16, 512] partial: rows + // [k_split*16, k_split*16+16) of `out_part_all`, column band at + // `n_op_inline64`. The splits touch disjoint regions, so TensorMap derives + // no edge between them and the kernel stores without an atomic. + const simpler::tmr::Tensor out_part_all_band_t12 = + out_part_all + .slice( + 0, static_cast(k_split_out_inline66 * 16), + static_cast(k_split_out_inline66 * 16) + 16 + ) + .slice( + 1, static_cast(n_op_inline64), static_cast(n_op_inline64) + 512 + ); + params_t12.add_output(out_part_all_band_t12); + params_t12.add_scalar(k_op_inline266); + params_t12.add_scalar(layer_hidden_base_inline151); + params_t12.add_scalar(n_op_inline64); + TaskId params_t12_deps[1]; + uint32_t params_t12_deps_count = 0; + if (out_proj_dummy_inline257.is_valid()) + params_t12_deps[params_t12_deps_count++] = out_proj_dummy_inline257; + params_t12.set_dependencies(params_t12_deps, params_t12_deps_count); + TaskOutputTensors task_12_outs = rt_submit_aic_task(13, params_t12); + TaskId out_tid_inline141 = task_12_outs.task_id(); + out_tids_inline271[out_idx_inline74] = out_tid_inline141; + } + + // Task 12r: partials_reduce_hidden_add, one per band the direct split-K tasks + // above touched. `out_proj_0`'s SPMD blocks accumulate into the same bands from + // a block_idx-derived offset, so the sum joins them atomically. A band holds + // only the splits the loop actually emitted — `out_idx` runs to + // N_OUT_DIRECT_inline61, so the last band can be short and the unwritten rows + // of its slab must stay out of the sum. + for (int64_t n_op_reduce = 0; n_op_reduce * 5 < N_OUT_DIRECT_inline61; n_op_reduce += 1) { + int64_t splits_here = N_OUT_DIRECT_inline61 - (n_op_reduce * 5); + if (splits_here > 5) splits_here = 5; + int64_t n_op_r = (n_op_reduce * 512); + CoreTaskArgs params_t12r; + const simpler::tmr::Tensor out_part_all_band_t12r = + out_part_all.slice(0, 0, static_cast(splits_here * 16)) + .slice(1, static_cast(n_op_r), static_cast(n_op_r) + 512); + params_t12r.add_input(out_part_all_band_t12r); + const simpler::tmr::Tensor attn_proj_fp32_inline220_band_t12r = attn_proj_fp32_inline220.slice( + 1, static_cast(n_op_r), static_cast(n_op_r) + 512 + ); + params_t12r.add_inout(attn_proj_fp32_inline220_band_t12r); + params_t12r.add_scalar(splits_here); + params_t12r.add_scalar(static_cast(512)); + TaskOutputTensors task_12r_outs = rt_submit_aiv_task(38, params_t12r); + (void)task_12r_outs; + } + + // Spmd out_proj_spmd: out_proj_0 + CoreTaskArgs params_t13; + params_t13.add_input(attn_out_inline282__ssa_v4); + params_t13.add_input(ext_wo); + // Bands 5..9 are this SPMD form's whole footprint, but its store addresses + // them absolutely: block b resolves flat index N_OUT_DIRECT + b to column + // (idx / 5) * 512 measured from the buffer base, not from the argument's base. + // A narrower view would shift that base and move every store out of range, so + // the declaration stays at the parent until the kernel indexes relatively. + params_t13.add_inout(attn_proj_fp32_inline220); + params_t13.add_scalar(N_OUT_DIRECT_inline61); + params_t13.add_scalar(layer_hidden_base_inline151); + params_t13.launch_spec.set_block_num(24); + TaskId params_t13_deps[1]; + uint32_t params_t13_deps_count = 0; + params_t13_deps[params_t13_deps_count++] = attn_done_tid_inline78; + params_t13.set_dependencies(params_t13_deps, params_t13_deps_count); + TaskOutputTensors task_13_outs = rt_submit_aic_task(14, params_t13); + TaskId out_proj_direct_tid_inline70 = task_13_outs.task_id(); + out_tids_inline271[N_OUT_DIRECT_inline61] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 1)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 2)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 3)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 4)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 5)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 6)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 7)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 8)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 9)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 10)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 11)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 12)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 13)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 14)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 15)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 16)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 17)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 18)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 19)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 20)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 21)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 22)] = out_proj_direct_tid_inline70; + out_tids_inline271[(N_OUT_DIRECT_inline61 + 23)] = out_proj_direct_tid_inline70; + int64_t k_base_inline111 = 0; + int64_t n_split_base_inline163 = 0; + TaskId _submit_deps_buf_inline165[10]; + for (int64_t __init_i = 0; __init_i < 10; ++__init_i) + _submit_deps_buf_inline165[__init_i] = TaskId::invalid(); + TaskId t__tmp_v39 = out_tids_inline271[(n_split_base_inline163 * 5)]; + _submit_deps_buf_inline165[0] = t__tmp_v39; + TaskId t__tmp_v40 = out_tids_inline271[((n_split_base_inline163 * 5) + 1)]; + _submit_deps_buf_inline165[1] = t__tmp_v40; + TaskId t__tmp_v41 = out_tids_inline271[((n_split_base_inline163 * 5) + 2)]; + _submit_deps_buf_inline165[2] = t__tmp_v41; + TaskId t__tmp_v42 = out_tids_inline271[((n_split_base_inline163 * 5) + 3)]; + _submit_deps_buf_inline165[3] = t__tmp_v42; + TaskId t__tmp_v43 = out_tids_inline271[((n_split_base_inline163 * 5) + 4)]; + _submit_deps_buf_inline165[4] = t__tmp_v43; + TaskId t__tmp_v44 = out_tids_inline271[((n_split_base_inline163 * 5) + 5)]; + _submit_deps_buf_inline165[5] = t__tmp_v44; + TaskId t__tmp_v45 = out_tids_inline271[((n_split_base_inline163 * 5) + 6)]; + _submit_deps_buf_inline165[6] = t__tmp_v45; + TaskId t__tmp_v46 = out_tids_inline271[((n_split_base_inline163 * 5) + 7)]; + _submit_deps_buf_inline165[7] = t__tmp_v46; + TaskId t__tmp_v47 = out_tids_inline271[((n_split_base_inline163 * 5) + 8)]; + _submit_deps_buf_inline165[8] = t__tmp_v47; + TaskId t__tmp_v48 = out_tids_inline271[((n_split_base_inline163 * 5) + 9)]; + _submit_deps_buf_inline165[9] = t__tmp_v48; + + // Task 14: residual_rms_cast + CoreTaskArgs params_t14; + // These five tasks write disjoint 1024-wide bands of `mlp_norm_in` and + // `post_norm_partial` with plain stores, one per k_base. Declaring the parents + // would order all five against each other even though they never overlap; the + // kernel's two stores index relative to these views. + const simpler::tmr::Tensor mlp_norm_in_inline71_band_t14 = mlp_norm_in_inline71.slice( + 1, static_cast(k_base_inline111), static_cast(k_base_inline111) + 1024 + ); + params_t14.add_inout(mlp_norm_in_inline71_band_t14); + const simpler::tmr::Tensor post_norm_partial_inline118_band_t14 = post_norm_partial_inline118.slice( + 1, static_cast(k_base_inline111), static_cast(k_base_inline111) + 1024 + ); + params_t14.add_inout(post_norm_partial_inline118_band_t14); + params_t14.add_input(attn_proj_fp32_inline220); + params_t14.add_input(cur__rv_v7); + params_t14.add_input(ext_post_rms_weight); + params_t14.add_scalar(k_base_inline111); + params_t14.add_scalar(i); + TaskId params_t14_deps[10]; + uint32_t params_t14_deps_count = 0; + if (_submit_deps_buf_inline165[0].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[0]; + if (_submit_deps_buf_inline165[1].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[1]; + if (_submit_deps_buf_inline165[2].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[2]; + if (_submit_deps_buf_inline165[3].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[3]; + if (_submit_deps_buf_inline165[4].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[4]; + if (_submit_deps_buf_inline165[5].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[5]; + if (_submit_deps_buf_inline165[6].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[6]; + if (_submit_deps_buf_inline165[7].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[7]; + if (_submit_deps_buf_inline165[8].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[8]; + if (_submit_deps_buf_inline165[9].is_valid()) + params_t14_deps[params_t14_deps_count++] = _submit_deps_buf_inline165[9]; + params_t14.set_dependencies(params_t14_deps, params_t14_deps_count); + params_t14.set_allow_early_resolve(true); + TaskOutputTensors task_14_outs = rt_submit_aiv_task(15, params_t14); + TaskId cast_tid_k_inline76 = task_14_outs.task_id(); + cast_tids_inline88[0] = cast_tid_k_inline76; + int64_t k_base_inline111__ssa_v1 = 1024; + int64_t n_split_base_inline163__ssa_v1 = 2; + TaskId _submit_deps_buf_inline165__ssa_v1[10]; + for (int64_t __init_i = 0; __init_i < 10; ++__init_i) + _submit_deps_buf_inline165__ssa_v1[__init_i] = TaskId::invalid(); + TaskId t__tmp_v51 = out_tids_inline271[(n_split_base_inline163__ssa_v1 * 5)]; + _submit_deps_buf_inline165__ssa_v1[0] = t__tmp_v51; + TaskId t__tmp_v52 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 1)]; + _submit_deps_buf_inline165__ssa_v1[1] = t__tmp_v52; + TaskId t__tmp_v53 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 2)]; + _submit_deps_buf_inline165__ssa_v1[2] = t__tmp_v53; + TaskId t__tmp_v54 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 3)]; + _submit_deps_buf_inline165__ssa_v1[3] = t__tmp_v54; + TaskId t__tmp_v55 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 4)]; + _submit_deps_buf_inline165__ssa_v1[4] = t__tmp_v55; + TaskId t__tmp_v56 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 5)]; + _submit_deps_buf_inline165__ssa_v1[5] = t__tmp_v56; + TaskId t__tmp_v57 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 6)]; + _submit_deps_buf_inline165__ssa_v1[6] = t__tmp_v57; + TaskId t__tmp_v58 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 7)]; + _submit_deps_buf_inline165__ssa_v1[7] = t__tmp_v58; + TaskId t__tmp_v59 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 8)]; + _submit_deps_buf_inline165__ssa_v1[8] = t__tmp_v59; + TaskId t__tmp_v60 = out_tids_inline271[((n_split_base_inline163__ssa_v1 * 5) + 9)]; + _submit_deps_buf_inline165__ssa_v1[9] = t__tmp_v60; + + // Task 15: residual_rms_cast_0 + CoreTaskArgs params_t15; + // These five tasks write disjoint 1024-wide bands of `mlp_norm_in` and + // `post_norm_partial` with plain stores, one per k_base. Declaring the parents + // would order all five against each other even though they never overlap; the + // kernel's two stores index relative to these views. + const simpler::tmr::Tensor mlp_norm_in_inline71_band_t15 = mlp_norm_in_inline71.slice( + 1, static_cast(k_base_inline111__ssa_v1), + static_cast(k_base_inline111__ssa_v1) + 1024 + ); + params_t15.add_inout(mlp_norm_in_inline71_band_t15); + const simpler::tmr::Tensor post_norm_partial_inline118_band_t15 = post_norm_partial_inline118.slice( + 1, static_cast(k_base_inline111__ssa_v1), + static_cast(k_base_inline111__ssa_v1) + 1024 + ); + params_t15.add_inout(post_norm_partial_inline118_band_t15); + params_t15.add_input(attn_proj_fp32_inline220); + params_t15.add_input(cur__rv_v7); + params_t15.add_input(ext_post_rms_weight); + params_t15.add_scalar(k_base_inline111__ssa_v1); + params_t15.add_scalar(i); + TaskId params_t15_deps[10]; + uint32_t params_t15_deps_count = 0; + if (_submit_deps_buf_inline165__ssa_v1[0].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[0]; + if (_submit_deps_buf_inline165__ssa_v1[1].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[1]; + if (_submit_deps_buf_inline165__ssa_v1[2].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[2]; + if (_submit_deps_buf_inline165__ssa_v1[3].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[3]; + if (_submit_deps_buf_inline165__ssa_v1[4].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[4]; + if (_submit_deps_buf_inline165__ssa_v1[5].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[5]; + if (_submit_deps_buf_inline165__ssa_v1[6].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[6]; + if (_submit_deps_buf_inline165__ssa_v1[7].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[7]; + if (_submit_deps_buf_inline165__ssa_v1[8].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[8]; + if (_submit_deps_buf_inline165__ssa_v1[9].is_valid()) + params_t15_deps[params_t15_deps_count++] = _submit_deps_buf_inline165__ssa_v1[9]; + params_t15.set_dependencies(params_t15_deps, params_t15_deps_count); + params_t15.set_allow_early_resolve(true); + TaskOutputTensors task_15_outs = rt_submit_aiv_task(16, params_t15); + TaskId cast_tid_k_inline76__ssa_v1 = task_15_outs.task_id(); + cast_tids_inline88[1] = cast_tid_k_inline76__ssa_v1; + int64_t k_base_inline111__ssa_v2 = 2048; + int64_t n_split_base_inline163__ssa_v2 = 4; + TaskId _submit_deps_buf_inline165__ssa_v2[10]; + for (int64_t __init_i = 0; __init_i < 10; ++__init_i) + _submit_deps_buf_inline165__ssa_v2[__init_i] = TaskId::invalid(); + TaskId t__tmp_v63 = out_tids_inline271[(n_split_base_inline163__ssa_v2 * 5)]; + _submit_deps_buf_inline165__ssa_v2[0] = t__tmp_v63; + TaskId t__tmp_v64 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 1)]; + _submit_deps_buf_inline165__ssa_v2[1] = t__tmp_v64; + TaskId t__tmp_v65 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 2)]; + _submit_deps_buf_inline165__ssa_v2[2] = t__tmp_v65; + TaskId t__tmp_v66 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 3)]; + _submit_deps_buf_inline165__ssa_v2[3] = t__tmp_v66; + TaskId t__tmp_v67 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 4)]; + _submit_deps_buf_inline165__ssa_v2[4] = t__tmp_v67; + TaskId t__tmp_v68 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 5)]; + _submit_deps_buf_inline165__ssa_v2[5] = t__tmp_v68; + TaskId t__tmp_v69 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 6)]; + _submit_deps_buf_inline165__ssa_v2[6] = t__tmp_v69; + TaskId t__tmp_v70 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 7)]; + _submit_deps_buf_inline165__ssa_v2[7] = t__tmp_v70; + TaskId t__tmp_v71 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 8)]; + _submit_deps_buf_inline165__ssa_v2[8] = t__tmp_v71; + TaskId t__tmp_v72 = out_tids_inline271[((n_split_base_inline163__ssa_v2 * 5) + 9)]; + _submit_deps_buf_inline165__ssa_v2[9] = t__tmp_v72; + + // Task 16: residual_rms_cast_1 + CoreTaskArgs params_t16; + // These five tasks write disjoint 1024-wide bands of `mlp_norm_in` and + // `post_norm_partial` with plain stores, one per k_base. Declaring the parents + // would order all five against each other even though they never overlap; the + // kernel's two stores index relative to these views. + const simpler::tmr::Tensor mlp_norm_in_inline71_band_t16 = mlp_norm_in_inline71.slice( + 1, static_cast(k_base_inline111__ssa_v2), + static_cast(k_base_inline111__ssa_v2) + 1024 + ); + params_t16.add_inout(mlp_norm_in_inline71_band_t16); + const simpler::tmr::Tensor post_norm_partial_inline118_band_t16 = post_norm_partial_inline118.slice( + 1, static_cast(k_base_inline111__ssa_v2), + static_cast(k_base_inline111__ssa_v2) + 1024 + ); + params_t16.add_inout(post_norm_partial_inline118_band_t16); + params_t16.add_input(attn_proj_fp32_inline220); + params_t16.add_input(cur__rv_v7); + params_t16.add_input(ext_post_rms_weight); + params_t16.add_scalar(k_base_inline111__ssa_v2); + params_t16.add_scalar(i); + TaskId params_t16_deps[10]; + uint32_t params_t16_deps_count = 0; + if (_submit_deps_buf_inline165__ssa_v2[0].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[0]; + if (_submit_deps_buf_inline165__ssa_v2[1].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[1]; + if (_submit_deps_buf_inline165__ssa_v2[2].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[2]; + if (_submit_deps_buf_inline165__ssa_v2[3].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[3]; + if (_submit_deps_buf_inline165__ssa_v2[4].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[4]; + if (_submit_deps_buf_inline165__ssa_v2[5].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[5]; + if (_submit_deps_buf_inline165__ssa_v2[6].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[6]; + if (_submit_deps_buf_inline165__ssa_v2[7].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[7]; + if (_submit_deps_buf_inline165__ssa_v2[8].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[8]; + if (_submit_deps_buf_inline165__ssa_v2[9].is_valid()) + params_t16_deps[params_t16_deps_count++] = _submit_deps_buf_inline165__ssa_v2[9]; + params_t16.set_dependencies(params_t16_deps, params_t16_deps_count); + params_t16.set_allow_early_resolve(true); + TaskOutputTensors task_16_outs = rt_submit_aiv_task(17, params_t16); + TaskId cast_tid_k_inline76__ssa_v2 = task_16_outs.task_id(); + cast_tids_inline88[2] = cast_tid_k_inline76__ssa_v2; + int64_t k_base_inline111__ssa_v3 = 3072; + int64_t n_split_base_inline163__ssa_v3 = 6; + TaskId _submit_deps_buf_inline165__ssa_v3[10]; + for (int64_t __init_i = 0; __init_i < 10; ++__init_i) + _submit_deps_buf_inline165__ssa_v3[__init_i] = TaskId::invalid(); + TaskId t__tmp_v75 = out_tids_inline271[(n_split_base_inline163__ssa_v3 * 5)]; + _submit_deps_buf_inline165__ssa_v3[0] = t__tmp_v75; + TaskId t__tmp_v76 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 1)]; + _submit_deps_buf_inline165__ssa_v3[1] = t__tmp_v76; + TaskId t__tmp_v77 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 2)]; + _submit_deps_buf_inline165__ssa_v3[2] = t__tmp_v77; + TaskId t__tmp_v78 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 3)]; + _submit_deps_buf_inline165__ssa_v3[3] = t__tmp_v78; + TaskId t__tmp_v79 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 4)]; + _submit_deps_buf_inline165__ssa_v3[4] = t__tmp_v79; + TaskId t__tmp_v80 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 5)]; + _submit_deps_buf_inline165__ssa_v3[5] = t__tmp_v80; + TaskId t__tmp_v81 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 6)]; + _submit_deps_buf_inline165__ssa_v3[6] = t__tmp_v81; + TaskId t__tmp_v82 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 7)]; + _submit_deps_buf_inline165__ssa_v3[7] = t__tmp_v82; + TaskId t__tmp_v83 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 8)]; + _submit_deps_buf_inline165__ssa_v3[8] = t__tmp_v83; + TaskId t__tmp_v84 = out_tids_inline271[((n_split_base_inline163__ssa_v3 * 5) + 9)]; + _submit_deps_buf_inline165__ssa_v3[9] = t__tmp_v84; + + // Task 17: residual_rms_cast_2 + CoreTaskArgs params_t17; + // These five tasks write disjoint 1024-wide bands of `mlp_norm_in` and + // `post_norm_partial` with plain stores, one per k_base. Declaring the parents + // would order all five against each other even though they never overlap; the + // kernel's two stores index relative to these views. + const simpler::tmr::Tensor mlp_norm_in_inline71_band_t17 = mlp_norm_in_inline71.slice( + 1, static_cast(k_base_inline111__ssa_v3), + static_cast(k_base_inline111__ssa_v3) + 1024 + ); + params_t17.add_inout(mlp_norm_in_inline71_band_t17); + const simpler::tmr::Tensor post_norm_partial_inline118_band_t17 = post_norm_partial_inline118.slice( + 1, static_cast(k_base_inline111__ssa_v3), + static_cast(k_base_inline111__ssa_v3) + 1024 + ); + params_t17.add_inout(post_norm_partial_inline118_band_t17); + params_t17.add_input(attn_proj_fp32_inline220); + params_t17.add_input(cur__rv_v7); + params_t17.add_input(ext_post_rms_weight); + params_t17.add_scalar(k_base_inline111__ssa_v3); + params_t17.add_scalar(i); + TaskId params_t17_deps[10]; + uint32_t params_t17_deps_count = 0; + if (_submit_deps_buf_inline165__ssa_v3[0].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[0]; + if (_submit_deps_buf_inline165__ssa_v3[1].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[1]; + if (_submit_deps_buf_inline165__ssa_v3[2].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[2]; + if (_submit_deps_buf_inline165__ssa_v3[3].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[3]; + if (_submit_deps_buf_inline165__ssa_v3[4].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[4]; + if (_submit_deps_buf_inline165__ssa_v3[5].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[5]; + if (_submit_deps_buf_inline165__ssa_v3[6].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[6]; + if (_submit_deps_buf_inline165__ssa_v3[7].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[7]; + if (_submit_deps_buf_inline165__ssa_v3[8].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[8]; + if (_submit_deps_buf_inline165__ssa_v3[9].is_valid()) + params_t17_deps[params_t17_deps_count++] = _submit_deps_buf_inline165__ssa_v3[9]; + params_t17.set_dependencies(params_t17_deps, params_t17_deps_count); + params_t17.set_allow_early_resolve(true); + TaskOutputTensors task_17_outs = rt_submit_aiv_task(18, params_t17); + TaskId cast_tid_k_inline76__ssa_v3 = task_17_outs.task_id(); + cast_tids_inline88[3] = cast_tid_k_inline76__ssa_v3; + int64_t k_base_inline111__ssa_v4 = 4096; + int64_t n_split_base_inline163__ssa_v4 = 8; + TaskId _submit_deps_buf_inline165__ssa_v4[10]; + for (int64_t __init_i = 0; __init_i < 10; ++__init_i) + _submit_deps_buf_inline165__ssa_v4[__init_i] = TaskId::invalid(); + TaskId t__tmp_v87 = out_tids_inline271[(n_split_base_inline163__ssa_v4 * 5)]; + _submit_deps_buf_inline165__ssa_v4[0] = t__tmp_v87; + TaskId t__tmp_v88 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 1)]; + _submit_deps_buf_inline165__ssa_v4[1] = t__tmp_v88; + TaskId t__tmp_v89 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 2)]; + _submit_deps_buf_inline165__ssa_v4[2] = t__tmp_v89; + TaskId t__tmp_v90 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 3)]; + _submit_deps_buf_inline165__ssa_v4[3] = t__tmp_v90; + TaskId t__tmp_v91 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 4)]; + _submit_deps_buf_inline165__ssa_v4[4] = t__tmp_v91; + TaskId t__tmp_v92 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 5)]; + _submit_deps_buf_inline165__ssa_v4[5] = t__tmp_v92; + TaskId t__tmp_v93 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 6)]; + _submit_deps_buf_inline165__ssa_v4[6] = t__tmp_v93; + TaskId t__tmp_v94 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 7)]; + _submit_deps_buf_inline165__ssa_v4[7] = t__tmp_v94; + TaskId t__tmp_v95 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 8)]; + _submit_deps_buf_inline165__ssa_v4[8] = t__tmp_v95; + TaskId t__tmp_v96 = out_tids_inline271[((n_split_base_inline163__ssa_v4 * 5) + 9)]; + _submit_deps_buf_inline165__ssa_v4[9] = t__tmp_v96; + + // Task 18: residual_rms_cast_3 + CoreTaskArgs params_t18; + // These five tasks write disjoint 1024-wide bands of `mlp_norm_in` and + // `post_norm_partial` with plain stores, one per k_base. Declaring the parents + // would order all five against each other even though they never overlap; the + // kernel's two stores index relative to these views. + const simpler::tmr::Tensor mlp_norm_in_inline71_band_t18 = mlp_norm_in_inline71.slice( + 1, static_cast(k_base_inline111__ssa_v4), + static_cast(k_base_inline111__ssa_v4) + 1024 + ); + params_t18.add_inout(mlp_norm_in_inline71_band_t18); + const simpler::tmr::Tensor post_norm_partial_inline118_band_t18 = post_norm_partial_inline118.slice( + 1, static_cast(k_base_inline111__ssa_v4), + static_cast(k_base_inline111__ssa_v4) + 1024 + ); + params_t18.add_inout(post_norm_partial_inline118_band_t18); + params_t18.add_input(attn_proj_fp32_inline220); + params_t18.add_input(cur__rv_v7); + params_t18.add_input(ext_post_rms_weight); + params_t18.add_scalar(k_base_inline111__ssa_v4); + params_t18.add_scalar(i); + TaskId params_t18_deps[10]; + uint32_t params_t18_deps_count = 0; + if (_submit_deps_buf_inline165__ssa_v4[0].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[0]; + if (_submit_deps_buf_inline165__ssa_v4[1].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[1]; + if (_submit_deps_buf_inline165__ssa_v4[2].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[2]; + if (_submit_deps_buf_inline165__ssa_v4[3].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[3]; + if (_submit_deps_buf_inline165__ssa_v4[4].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[4]; + if (_submit_deps_buf_inline165__ssa_v4[5].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[5]; + if (_submit_deps_buf_inline165__ssa_v4[6].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[6]; + if (_submit_deps_buf_inline165__ssa_v4[7].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[7]; + if (_submit_deps_buf_inline165__ssa_v4[8].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[8]; + if (_submit_deps_buf_inline165__ssa_v4[9].is_valid()) + params_t18_deps[params_t18_deps_count++] = _submit_deps_buf_inline165__ssa_v4[9]; + params_t18.set_dependencies(params_t18_deps, params_t18_deps_count); + params_t18.set_allow_early_resolve(true); + TaskOutputTensors task_18_outs = rt_submit_aiv_task(19, params_t18); + TaskId cast_tid_k_inline76__ssa_v4 = task_18_outs.task_id(); + cast_tids_inline88[4] = cast_tid_k_inline76__ssa_v4; + + // Task 19: post_rms_reduce + CoreTaskArgs params_t19; + params_t19.add_input(attn_proj_fp32_inline220); + params_t19.add_input(cur__rv_v7); + params_t19.add_inout(inv_rms_tile_inline126); + TaskId params_t19_deps[50]; + uint32_t params_t19_deps_count = 0; + if (out_tids_inline271[0].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[0]; + if (out_tids_inline271[1].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[1]; + if (out_tids_inline271[2].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[2]; + if (out_tids_inline271[3].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[3]; + if (out_tids_inline271[4].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[4]; + if (out_tids_inline271[5].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[5]; + if (out_tids_inline271[6].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[6]; + if (out_tids_inline271[7].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[7]; + if (out_tids_inline271[8].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[8]; + if (out_tids_inline271[9].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[9]; + if (out_tids_inline271[10].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[10]; + if (out_tids_inline271[11].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[11]; + if (out_tids_inline271[12].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[12]; + if (out_tids_inline271[13].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[13]; + if (out_tids_inline271[14].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[14]; + if (out_tids_inline271[15].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[15]; + if (out_tids_inline271[16].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[16]; + if (out_tids_inline271[17].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[17]; + if (out_tids_inline271[18].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[18]; + if (out_tids_inline271[19].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[19]; + if (out_tids_inline271[20].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[20]; + if (out_tids_inline271[21].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[21]; + if (out_tids_inline271[22].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[22]; + if (out_tids_inline271[23].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[23]; + if (out_tids_inline271[24].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[24]; + if (out_tids_inline271[25].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[25]; + if (out_tids_inline271[26].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[26]; + if (out_tids_inline271[27].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[27]; + if (out_tids_inline271[28].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[28]; + if (out_tids_inline271[29].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[29]; + if (out_tids_inline271[30].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[30]; + if (out_tids_inline271[31].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[31]; + if (out_tids_inline271[32].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[32]; + if (out_tids_inline271[33].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[33]; + if (out_tids_inline271[34].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[34]; + if (out_tids_inline271[35].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[35]; + if (out_tids_inline271[36].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[36]; + if (out_tids_inline271[37].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[37]; + if (out_tids_inline271[38].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[38]; + if (out_tids_inline271[39].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[39]; + if (out_tids_inline271[40].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[40]; + if (out_tids_inline271[41].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[41]; + if (out_tids_inline271[42].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[42]; + if (out_tids_inline271[43].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[43]; + if (out_tids_inline271[44].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[44]; + if (out_tids_inline271[45].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[45]; + if (out_tids_inline271[46].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[46]; + if (out_tids_inline271[47].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[47]; + if (out_tids_inline271[48].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[48]; + if (out_tids_inline271[49].is_valid()) + params_t19_deps[params_t19_deps_count++] = out_tids_inline271[49]; + params_t19.set_dependencies(params_t19_deps, params_t19_deps_count); + TaskOutputTensors task_19_outs = rt_submit_aiv_task(20, params_t19); + TaskId reduce_tid_inline226 = task_19_outs.task_id(); + int64_t gu_k0_inline131 = 0; + TaskId _submit_deps_buf_inline236[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline236[__init_i] = TaskId::invalid(); + TaskId t__tmp_v105 = cast_tids_inline88[0]; + _submit_deps_buf_inline236[0] = t__tmp_v105; + + // Phase-fence barrier 3: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_3; + TaskId params_phase_fence_barrier_3_deps[1]; + uint32_t params_phase_fence_barrier_3_deps_count = 0; + if (_submit_deps_buf_inline236[0].is_valid()) + params_phase_fence_barrier_3_deps[params_phase_fence_barrier_3_deps_count++] = + _submit_deps_buf_inline236[0]; + params_phase_fence_barrier_3.set_dependencies( + params_phase_fence_barrier_3_deps, params_phase_fence_barrier_3_deps_count + ); + TaskId t__tmp_v106 = TaskId::invalid(); + if (params_phase_fence_barrier_3_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_3_outs = + rt_submit_dummy_task(params_phase_fence_barrier_3); + t__tmp_v106 = phase_fence_barrier_3_outs.task_id(); + } + gate_late_tids_inline249[0] = t__tmp_v106; + TaskId _submit_deps_buf_inline225[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline225[__init_i] = TaskId::invalid(); + TaskId t__tmp_v107 = cast_tids_inline88[0]; + _submit_deps_buf_inline225[0] = t__tmp_v107; + + // Phase-fence barrier 4: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_4; + TaskId params_phase_fence_barrier_4_deps[1]; + uint32_t params_phase_fence_barrier_4_deps_count = 0; + if (_submit_deps_buf_inline225[0].is_valid()) + params_phase_fence_barrier_4_deps[params_phase_fence_barrier_4_deps_count++] = + _submit_deps_buf_inline225[0]; + params_phase_fence_barrier_4.set_dependencies( + params_phase_fence_barrier_4_deps, params_phase_fence_barrier_4_deps_count + ); + TaskId t__tmp_v108 = TaskId::invalid(); + if (params_phase_fence_barrier_4_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_4_outs = + rt_submit_dummy_task(params_phase_fence_barrier_4); + t__tmp_v108 = phase_fence_barrier_4_outs.task_id(); + } + up_late_tids_inline69[0] = t__tmp_v108; + TaskId _submit_deps_buf_inline237[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline237[__init_i] = TaskId::invalid(); + TaskId t__tmp_v109 = cast_tids_inline88[0]; + _submit_deps_buf_inline237[0] = t__tmp_v109; + + // Spmd gate_proj_spmd: gate_proj + CoreTaskArgs params_t20; + params_t20.add_input(mlp_norm_in_inline71); + params_t20.add_input(ext_w_gate); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor gate_acc_all_inline203_spmd_t20 = + gate_acc_all_inline203.slice(1, 0, 6144); + params_t20.add_inout(gate_acc_all_inline203_spmd_t20); + params_t20.add_scalar(gu_k0_inline131); + params_t20.add_scalar(layer_hidden_base_inline151); + params_t20.launch_spec.set_block_num(6); + TaskId params_t20_deps[1]; + uint32_t params_t20_deps_count = 0; + if (_submit_deps_buf_inline237[0].is_valid()) + params_t20_deps[params_t20_deps_count++] = _submit_deps_buf_inline237[0]; + params_t20.set_dependencies(params_t20_deps, params_t20_deps_count); + TaskOutputTensors task_20_outs = rt_submit_aic_task(21, params_t20); + TaskId gate_spmd_tid_inline245 = task_20_outs.task_id(); + gate_tids_inline56[0] = gate_spmd_tid_inline245; + gate_tids_inline56[5] = gate_spmd_tid_inline245; + gate_tids_inline56[10] = gate_spmd_tid_inline245; + gate_tids_inline56[15] = gate_spmd_tid_inline245; + gate_tids_inline56[20] = gate_spmd_tid_inline245; + gate_tids_inline56[25] = gate_spmd_tid_inline245; + TaskId _submit_deps_buf_inline260[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline260[__init_i] = TaskId::invalid(); + TaskId t__tmp_v110 = cast_tids_inline88[0]; + _submit_deps_buf_inline260[0] = t__tmp_v110; + + // Spmd up_proj_spmd: up_proj + CoreTaskArgs params_t21; + params_t21.add_input(mlp_norm_in_inline71); + params_t21.add_input(ext_w_up); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor up_acc_all_inline303_spmd_t21 = up_acc_all_inline303.slice(1, 0, 6144); + params_t21.add_inout(up_acc_all_inline303_spmd_t21); + params_t21.add_scalar(gu_k0_inline131); + params_t21.add_scalar(layer_hidden_base_inline151); + params_t21.launch_spec.set_block_num(6); + TaskId params_t21_deps[1]; + uint32_t params_t21_deps_count = 0; + if (_submit_deps_buf_inline260[0].is_valid()) + params_t21_deps[params_t21_deps_count++] = _submit_deps_buf_inline260[0]; + params_t21.set_dependencies(params_t21_deps, params_t21_deps_count); + TaskOutputTensors task_21_outs = rt_submit_aic_task(22, params_t21); + TaskId up_spmd_tid_inline264 = task_21_outs.task_id(); + up_tids_inline310[0] = up_spmd_tid_inline264; + up_tids_inline310[5] = up_spmd_tid_inline264; + up_tids_inline310[10] = up_spmd_tid_inline264; + up_tids_inline310[15] = up_spmd_tid_inline264; + up_tids_inline310[20] = up_spmd_tid_inline264; + up_tids_inline310[25] = up_spmd_tid_inline264; + int64_t gu_k0_inline131__ssa_v1 = 1024; + TaskId _submit_deps_buf_inline236__ssa_v1[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline236__ssa_v1[__init_i] = TaskId::invalid(); + TaskId t__tmp_v111 = cast_tids_inline88[1]; + _submit_deps_buf_inline236__ssa_v1[0] = t__tmp_v111; + + // Phase-fence barrier 5: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_5; + TaskId params_phase_fence_barrier_5_deps[1]; + uint32_t params_phase_fence_barrier_5_deps_count = 0; + if (_submit_deps_buf_inline236__ssa_v1[0].is_valid()) + params_phase_fence_barrier_5_deps[params_phase_fence_barrier_5_deps_count++] = + _submit_deps_buf_inline236__ssa_v1[0]; + params_phase_fence_barrier_5.set_dependencies( + params_phase_fence_barrier_5_deps, params_phase_fence_barrier_5_deps_count + ); + TaskId t__tmp_v112 = TaskId::invalid(); + if (params_phase_fence_barrier_5_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_5_outs = + rt_submit_dummy_task(params_phase_fence_barrier_5); + t__tmp_v112 = phase_fence_barrier_5_outs.task_id(); + } + gate_late_tids_inline249[1] = t__tmp_v112; + TaskId _submit_deps_buf_inline225__ssa_v1[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline225__ssa_v1[__init_i] = TaskId::invalid(); + TaskId t__tmp_v113 = cast_tids_inline88[1]; + _submit_deps_buf_inline225__ssa_v1[0] = t__tmp_v113; + + // Phase-fence barrier 6: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_6; + TaskId params_phase_fence_barrier_6_deps[1]; + uint32_t params_phase_fence_barrier_6_deps_count = 0; + if (_submit_deps_buf_inline225__ssa_v1[0].is_valid()) + params_phase_fence_barrier_6_deps[params_phase_fence_barrier_6_deps_count++] = + _submit_deps_buf_inline225__ssa_v1[0]; + params_phase_fence_barrier_6.set_dependencies( + params_phase_fence_barrier_6_deps, params_phase_fence_barrier_6_deps_count + ); + TaskId t__tmp_v114 = TaskId::invalid(); + if (params_phase_fence_barrier_6_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_6_outs = + rt_submit_dummy_task(params_phase_fence_barrier_6); + t__tmp_v114 = phase_fence_barrier_6_outs.task_id(); + } + up_late_tids_inline69[1] = t__tmp_v114; + TaskId _submit_deps_buf_inline237__ssa_v1[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline237__ssa_v1[__init_i] = TaskId::invalid(); + TaskId t__tmp_v115 = cast_tids_inline88[1]; + _submit_deps_buf_inline237__ssa_v1[0] = t__tmp_v115; + + // Spmd gate_proj_spmd_0: gate_proj_0 + CoreTaskArgs params_t22; + params_t22.add_input(mlp_norm_in_inline71); + params_t22.add_input(ext_w_gate); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor gate_acc_all_inline203_spmd_t22 = + gate_acc_all_inline203.slice(1, 0, 6144); + params_t22.add_inout(gate_acc_all_inline203_spmd_t22); + params_t22.add_scalar(gu_k0_inline131__ssa_v1); + params_t22.add_scalar(layer_hidden_base_inline151); + params_t22.launch_spec.set_block_num(6); + TaskId params_t22_deps[1]; + uint32_t params_t22_deps_count = 0; + if (_submit_deps_buf_inline237__ssa_v1[0].is_valid()) + params_t22_deps[params_t22_deps_count++] = _submit_deps_buf_inline237__ssa_v1[0]; + params_t22.set_dependencies(params_t22_deps, params_t22_deps_count); + TaskOutputTensors task_22_outs = rt_submit_aic_task(23, params_t22); + TaskId gate_spmd_tid_inline245__ssa_v1 = task_22_outs.task_id(); + gate_tids_inline56[1] = gate_spmd_tid_inline245__ssa_v1; + gate_tids_inline56[6] = gate_spmd_tid_inline245__ssa_v1; + gate_tids_inline56[11] = gate_spmd_tid_inline245__ssa_v1; + gate_tids_inline56[16] = gate_spmd_tid_inline245__ssa_v1; + gate_tids_inline56[21] = gate_spmd_tid_inline245__ssa_v1; + gate_tids_inline56[26] = gate_spmd_tid_inline245__ssa_v1; + TaskId _submit_deps_buf_inline260__ssa_v1[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline260__ssa_v1[__init_i] = TaskId::invalid(); + TaskId t__tmp_v116 = cast_tids_inline88[1]; + _submit_deps_buf_inline260__ssa_v1[0] = t__tmp_v116; + + // Spmd up_proj_spmd_0: up_proj_0 + CoreTaskArgs params_t23; + params_t23.add_input(mlp_norm_in_inline71); + params_t23.add_input(ext_w_up); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor up_acc_all_inline303_spmd_t23 = up_acc_all_inline303.slice(1, 0, 6144); + params_t23.add_inout(up_acc_all_inline303_spmd_t23); + params_t23.add_scalar(gu_k0_inline131__ssa_v1); + params_t23.add_scalar(layer_hidden_base_inline151); + params_t23.launch_spec.set_block_num(6); + TaskId params_t23_deps[1]; + uint32_t params_t23_deps_count = 0; + if (_submit_deps_buf_inline260__ssa_v1[0].is_valid()) + params_t23_deps[params_t23_deps_count++] = _submit_deps_buf_inline260__ssa_v1[0]; + params_t23.set_dependencies(params_t23_deps, params_t23_deps_count); + TaskOutputTensors task_23_outs = rt_submit_aic_task(24, params_t23); + TaskId up_spmd_tid_inline264__ssa_v1 = task_23_outs.task_id(); + up_tids_inline310[1] = up_spmd_tid_inline264__ssa_v1; + up_tids_inline310[6] = up_spmd_tid_inline264__ssa_v1; + up_tids_inline310[11] = up_spmd_tid_inline264__ssa_v1; + up_tids_inline310[16] = up_spmd_tid_inline264__ssa_v1; + up_tids_inline310[21] = up_spmd_tid_inline264__ssa_v1; + up_tids_inline310[26] = up_spmd_tid_inline264__ssa_v1; + int64_t gu_k0_inline131__ssa_v2 = 2048; + TaskId _submit_deps_buf_inline236__ssa_v2[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline236__ssa_v2[__init_i] = TaskId::invalid(); + TaskId t__tmp_v117 = cast_tids_inline88[2]; + _submit_deps_buf_inline236__ssa_v2[0] = t__tmp_v117; + + // Phase-fence barrier 7: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_7; + TaskId params_phase_fence_barrier_7_deps[1]; + uint32_t params_phase_fence_barrier_7_deps_count = 0; + if (_submit_deps_buf_inline236__ssa_v2[0].is_valid()) + params_phase_fence_barrier_7_deps[params_phase_fence_barrier_7_deps_count++] = + _submit_deps_buf_inline236__ssa_v2[0]; + params_phase_fence_barrier_7.set_dependencies( + params_phase_fence_barrier_7_deps, params_phase_fence_barrier_7_deps_count + ); + TaskId t__tmp_v118 = TaskId::invalid(); + if (params_phase_fence_barrier_7_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_7_outs = + rt_submit_dummy_task(params_phase_fence_barrier_7); + t__tmp_v118 = phase_fence_barrier_7_outs.task_id(); + } + gate_late_tids_inline249[2] = t__tmp_v118; + TaskId _submit_deps_buf_inline225__ssa_v2[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline225__ssa_v2[__init_i] = TaskId::invalid(); + TaskId t__tmp_v119 = cast_tids_inline88[2]; + _submit_deps_buf_inline225__ssa_v2[0] = t__tmp_v119; + + // Phase-fence barrier 8: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_8; + TaskId params_phase_fence_barrier_8_deps[1]; + uint32_t params_phase_fence_barrier_8_deps_count = 0; + if (_submit_deps_buf_inline225__ssa_v2[0].is_valid()) + params_phase_fence_barrier_8_deps[params_phase_fence_barrier_8_deps_count++] = + _submit_deps_buf_inline225__ssa_v2[0]; + params_phase_fence_barrier_8.set_dependencies( + params_phase_fence_barrier_8_deps, params_phase_fence_barrier_8_deps_count + ); + TaskId t__tmp_v120 = TaskId::invalid(); + if (params_phase_fence_barrier_8_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_8_outs = + rt_submit_dummy_task(params_phase_fence_barrier_8); + t__tmp_v120 = phase_fence_barrier_8_outs.task_id(); + } + up_late_tids_inline69[2] = t__tmp_v120; + TaskId _submit_deps_buf_inline237__ssa_v2[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline237__ssa_v2[__init_i] = TaskId::invalid(); + TaskId t__tmp_v121 = cast_tids_inline88[2]; + _submit_deps_buf_inline237__ssa_v2[0] = t__tmp_v121; + + // Spmd gate_proj_spmd_1: gate_proj_1 + CoreTaskArgs params_t24; + params_t24.add_input(mlp_norm_in_inline71); + params_t24.add_input(ext_w_gate); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor gate_acc_all_inline203_spmd_t24 = + gate_acc_all_inline203.slice(1, 0, 6144); + params_t24.add_inout(gate_acc_all_inline203_spmd_t24); + params_t24.add_scalar(gu_k0_inline131__ssa_v2); + params_t24.add_scalar(layer_hidden_base_inline151); + params_t24.launch_spec.set_block_num(6); + TaskId params_t24_deps[1]; + uint32_t params_t24_deps_count = 0; + if (_submit_deps_buf_inline237__ssa_v2[0].is_valid()) + params_t24_deps[params_t24_deps_count++] = _submit_deps_buf_inline237__ssa_v2[0]; + params_t24.set_dependencies(params_t24_deps, params_t24_deps_count); + TaskOutputTensors task_24_outs = rt_submit_aic_task(25, params_t24); + TaskId gate_spmd_tid_inline245__ssa_v2 = task_24_outs.task_id(); + gate_tids_inline56[2] = gate_spmd_tid_inline245__ssa_v2; + gate_tids_inline56[7] = gate_spmd_tid_inline245__ssa_v2; + gate_tids_inline56[12] = gate_spmd_tid_inline245__ssa_v2; + gate_tids_inline56[17] = gate_spmd_tid_inline245__ssa_v2; + gate_tids_inline56[22] = gate_spmd_tid_inline245__ssa_v2; + gate_tids_inline56[27] = gate_spmd_tid_inline245__ssa_v2; + TaskId _submit_deps_buf_inline260__ssa_v2[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline260__ssa_v2[__init_i] = TaskId::invalid(); + TaskId t__tmp_v122 = cast_tids_inline88[2]; + _submit_deps_buf_inline260__ssa_v2[0] = t__tmp_v122; + + // Spmd up_proj_spmd_1: up_proj_1 + CoreTaskArgs params_t25; + params_t25.add_input(mlp_norm_in_inline71); + params_t25.add_input(ext_w_up); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor up_acc_all_inline303_spmd_t25 = up_acc_all_inline303.slice(1, 0, 6144); + params_t25.add_inout(up_acc_all_inline303_spmd_t25); + params_t25.add_scalar(gu_k0_inline131__ssa_v2); + params_t25.add_scalar(layer_hidden_base_inline151); + params_t25.launch_spec.set_block_num(6); + TaskId params_t25_deps[1]; + uint32_t params_t25_deps_count = 0; + if (_submit_deps_buf_inline260__ssa_v2[0].is_valid()) + params_t25_deps[params_t25_deps_count++] = _submit_deps_buf_inline260__ssa_v2[0]; + params_t25.set_dependencies(params_t25_deps, params_t25_deps_count); + TaskOutputTensors task_25_outs = rt_submit_aic_task(26, params_t25); + TaskId up_spmd_tid_inline264__ssa_v2 = task_25_outs.task_id(); + up_tids_inline310[2] = up_spmd_tid_inline264__ssa_v2; + up_tids_inline310[7] = up_spmd_tid_inline264__ssa_v2; + up_tids_inline310[12] = up_spmd_tid_inline264__ssa_v2; + up_tids_inline310[17] = up_spmd_tid_inline264__ssa_v2; + up_tids_inline310[22] = up_spmd_tid_inline264__ssa_v2; + up_tids_inline310[27] = up_spmd_tid_inline264__ssa_v2; + int64_t gu_k0_inline131__ssa_v3 = 3072; + TaskId _submit_deps_buf_inline236__ssa_v3[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline236__ssa_v3[__init_i] = TaskId::invalid(); + TaskId t__tmp_v123 = cast_tids_inline88[3]; + _submit_deps_buf_inline236__ssa_v3[0] = t__tmp_v123; + + // Phase-fence barrier 9: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_9; + TaskId params_phase_fence_barrier_9_deps[1]; + uint32_t params_phase_fence_barrier_9_deps_count = 0; + if (_submit_deps_buf_inline236__ssa_v3[0].is_valid()) + params_phase_fence_barrier_9_deps[params_phase_fence_barrier_9_deps_count++] = + _submit_deps_buf_inline236__ssa_v3[0]; + params_phase_fence_barrier_9.set_dependencies( + params_phase_fence_barrier_9_deps, params_phase_fence_barrier_9_deps_count + ); + TaskId t__tmp_v124 = TaskId::invalid(); + if (params_phase_fence_barrier_9_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_9_outs = + rt_submit_dummy_task(params_phase_fence_barrier_9); + t__tmp_v124 = phase_fence_barrier_9_outs.task_id(); + } + gate_late_tids_inline249[3] = t__tmp_v124; + TaskId _submit_deps_buf_inline225__ssa_v3[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline225__ssa_v3[__init_i] = TaskId::invalid(); + TaskId t__tmp_v125 = cast_tids_inline88[3]; + _submit_deps_buf_inline225__ssa_v3[0] = t__tmp_v125; + + // Phase-fence barrier 10: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_10; + TaskId params_phase_fence_barrier_10_deps[1]; + uint32_t params_phase_fence_barrier_10_deps_count = 0; + if (_submit_deps_buf_inline225__ssa_v3[0].is_valid()) + params_phase_fence_barrier_10_deps[params_phase_fence_barrier_10_deps_count++] = + _submit_deps_buf_inline225__ssa_v3[0]; + params_phase_fence_barrier_10.set_dependencies( + params_phase_fence_barrier_10_deps, params_phase_fence_barrier_10_deps_count + ); + TaskId t__tmp_v126 = TaskId::invalid(); + if (params_phase_fence_barrier_10_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_10_outs = + rt_submit_dummy_task(params_phase_fence_barrier_10); + t__tmp_v126 = phase_fence_barrier_10_outs.task_id(); + } + up_late_tids_inline69[3] = t__tmp_v126; + TaskId _submit_deps_buf_inline237__ssa_v3[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline237__ssa_v3[__init_i] = TaskId::invalid(); + TaskId t__tmp_v127 = cast_tids_inline88[3]; + _submit_deps_buf_inline237__ssa_v3[0] = t__tmp_v127; + + // Spmd gate_proj_spmd_2: gate_proj_2 + CoreTaskArgs params_t26; + params_t26.add_input(mlp_norm_in_inline71); + params_t26.add_input(ext_w_gate); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor gate_acc_all_inline203_spmd_t26 = + gate_acc_all_inline203.slice(1, 0, 6144); + params_t26.add_inout(gate_acc_all_inline203_spmd_t26); + params_t26.add_scalar(gu_k0_inline131__ssa_v3); + params_t26.add_scalar(layer_hidden_base_inline151); + params_t26.launch_spec.set_block_num(6); + TaskId params_t26_deps[1]; + uint32_t params_t26_deps_count = 0; + if (_submit_deps_buf_inline237__ssa_v3[0].is_valid()) + params_t26_deps[params_t26_deps_count++] = _submit_deps_buf_inline237__ssa_v3[0]; + params_t26.set_dependencies(params_t26_deps, params_t26_deps_count); + TaskOutputTensors task_26_outs = rt_submit_aic_task(27, params_t26); + TaskId gate_spmd_tid_inline245__ssa_v3 = task_26_outs.task_id(); + gate_tids_inline56[3] = gate_spmd_tid_inline245__ssa_v3; + gate_tids_inline56[8] = gate_spmd_tid_inline245__ssa_v3; + gate_tids_inline56[13] = gate_spmd_tid_inline245__ssa_v3; + gate_tids_inline56[18] = gate_spmd_tid_inline245__ssa_v3; + gate_tids_inline56[23] = gate_spmd_tid_inline245__ssa_v3; + gate_tids_inline56[28] = gate_spmd_tid_inline245__ssa_v3; + TaskId _submit_deps_buf_inline260__ssa_v3[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline260__ssa_v3[__init_i] = TaskId::invalid(); + TaskId t__tmp_v128 = cast_tids_inline88[3]; + _submit_deps_buf_inline260__ssa_v3[0] = t__tmp_v128; + + // Spmd up_proj_spmd_2: up_proj_2 + CoreTaskArgs params_t27; + params_t27.add_input(mlp_norm_in_inline71); + params_t27.add_input(ext_w_up); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor up_acc_all_inline303_spmd_t27 = up_acc_all_inline303.slice(1, 0, 6144); + params_t27.add_inout(up_acc_all_inline303_spmd_t27); + params_t27.add_scalar(gu_k0_inline131__ssa_v3); + params_t27.add_scalar(layer_hidden_base_inline151); + params_t27.launch_spec.set_block_num(6); + TaskId params_t27_deps[1]; + uint32_t params_t27_deps_count = 0; + if (_submit_deps_buf_inline260__ssa_v3[0].is_valid()) + params_t27_deps[params_t27_deps_count++] = _submit_deps_buf_inline260__ssa_v3[0]; + params_t27.set_dependencies(params_t27_deps, params_t27_deps_count); + TaskOutputTensors task_27_outs = rt_submit_aic_task(28, params_t27); + TaskId up_spmd_tid_inline264__ssa_v3 = task_27_outs.task_id(); + up_tids_inline310[3] = up_spmd_tid_inline264__ssa_v3; + up_tids_inline310[8] = up_spmd_tid_inline264__ssa_v3; + up_tids_inline310[13] = up_spmd_tid_inline264__ssa_v3; + up_tids_inline310[18] = up_spmd_tid_inline264__ssa_v3; + up_tids_inline310[23] = up_spmd_tid_inline264__ssa_v3; + up_tids_inline310[28] = up_spmd_tid_inline264__ssa_v3; + int64_t gu_k0_inline131__ssa_v4 = 4096; + TaskId _submit_deps_buf_inline236__ssa_v4[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline236__ssa_v4[__init_i] = TaskId::invalid(); + TaskId t__tmp_v129 = cast_tids_inline88[4]; + _submit_deps_buf_inline236__ssa_v4[0] = t__tmp_v129; + + // Phase-fence barrier 11: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_11; + TaskId params_phase_fence_barrier_11_deps[1]; + uint32_t params_phase_fence_barrier_11_deps_count = 0; + if (_submit_deps_buf_inline236__ssa_v4[0].is_valid()) + params_phase_fence_barrier_11_deps[params_phase_fence_barrier_11_deps_count++] = + _submit_deps_buf_inline236__ssa_v4[0]; + params_phase_fence_barrier_11.set_dependencies( + params_phase_fence_barrier_11_deps, params_phase_fence_barrier_11_deps_count + ); + TaskId t__tmp_v130 = TaskId::invalid(); + if (params_phase_fence_barrier_11_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_11_outs = + rt_submit_dummy_task(params_phase_fence_barrier_11); + t__tmp_v130 = phase_fence_barrier_11_outs.task_id(); + } + gate_late_tids_inline249[4] = t__tmp_v130; + TaskId _submit_deps_buf_inline225__ssa_v4[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline225__ssa_v4[__init_i] = TaskId::invalid(); + TaskId t__tmp_v131 = cast_tids_inline88[4]; + _submit_deps_buf_inline225__ssa_v4[0] = t__tmp_v131; + + // Phase-fence barrier 12: dependency-only dummy task + CoreTaskArgs params_phase_fence_barrier_12; + TaskId params_phase_fence_barrier_12_deps[1]; + uint32_t params_phase_fence_barrier_12_deps_count = 0; + if (_submit_deps_buf_inline225__ssa_v4[0].is_valid()) + params_phase_fence_barrier_12_deps[params_phase_fence_barrier_12_deps_count++] = + _submit_deps_buf_inline225__ssa_v4[0]; + params_phase_fence_barrier_12.set_dependencies( + params_phase_fence_barrier_12_deps, params_phase_fence_barrier_12_deps_count + ); + TaskId t__tmp_v132 = TaskId::invalid(); + if (params_phase_fence_barrier_12_deps_count > 0) { + TaskOutputTensors phase_fence_barrier_12_outs = + rt_submit_dummy_task(params_phase_fence_barrier_12); + t__tmp_v132 = phase_fence_barrier_12_outs.task_id(); + } + up_late_tids_inline69[4] = t__tmp_v132; + TaskId _submit_deps_buf_inline237__ssa_v4[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline237__ssa_v4[__init_i] = TaskId::invalid(); + TaskId t__tmp_v133 = cast_tids_inline88[4]; + _submit_deps_buf_inline237__ssa_v4[0] = t__tmp_v133; + + // Spmd gate_proj_spmd_3: gate_proj_3 + CoreTaskArgs params_t28; + params_t28.add_input(mlp_norm_in_inline71); + params_t28.add_input(ext_w_gate); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor gate_acc_all_inline203_spmd_t28 = + gate_acc_all_inline203.slice(1, 0, 6144); + params_t28.add_inout(gate_acc_all_inline203_spmd_t28); + params_t28.add_scalar(gu_k0_inline131__ssa_v4); + params_t28.add_scalar(layer_hidden_base_inline151); + params_t28.launch_spec.set_block_num(6); + TaskId params_t28_deps[1]; + uint32_t params_t28_deps_count = 0; + if (_submit_deps_buf_inline237__ssa_v4[0].is_valid()) + params_t28_deps[params_t28_deps_count++] = _submit_deps_buf_inline237__ssa_v4[0]; + params_t28.set_dependencies(params_t28_deps, params_t28_deps_count); + TaskOutputTensors task_28_outs = rt_submit_aic_task(29, params_t28); + TaskId gate_spmd_tid_inline245__ssa_v4 = task_28_outs.task_id(); + gate_tids_inline56[4] = gate_spmd_tid_inline245__ssa_v4; + gate_tids_inline56[9] = gate_spmd_tid_inline245__ssa_v4; + gate_tids_inline56[14] = gate_spmd_tid_inline245__ssa_v4; + gate_tids_inline56[19] = gate_spmd_tid_inline245__ssa_v4; + gate_tids_inline56[24] = gate_spmd_tid_inline245__ssa_v4; + gate_tids_inline56[29] = gate_spmd_tid_inline245__ssa_v4; + TaskId _submit_deps_buf_inline260__ssa_v4[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline260__ssa_v4[__init_i] = TaskId::invalid(); + TaskId t__tmp_v134 = cast_tids_inline88[4]; + _submit_deps_buf_inline260__ssa_v4[0] = t__tmp_v134; + + // Spmd up_proj_spmd_3: up_proj_3 + CoreTaskArgs params_t29; + params_t29.add_input(mlp_norm_in_inline71); + params_t29.add_input(ext_w_up); + // The SPMD form covers bands 0..5: its 6 blocks each store one 1024-wide + // band at `block_idx * 1024`, so [0, 6144) is the whole footprint. Declaring + // the parent would order it against every banded write on bands 6..16. + const simpler::tmr::Tensor up_acc_all_inline303_spmd_t29 = up_acc_all_inline303.slice(1, 0, 6144); + params_t29.add_inout(up_acc_all_inline303_spmd_t29); + params_t29.add_scalar(gu_k0_inline131__ssa_v4); + params_t29.add_scalar(layer_hidden_base_inline151); + params_t29.launch_spec.set_block_num(6); + TaskId params_t29_deps[1]; + uint32_t params_t29_deps_count = 0; + if (_submit_deps_buf_inline260__ssa_v4[0].is_valid()) + params_t29_deps[params_t29_deps_count++] = _submit_deps_buf_inline260__ssa_v4[0]; + params_t29.set_dependencies(params_t29_deps, params_t29_deps_count); + TaskOutputTensors task_29_outs = rt_submit_aic_task(30, params_t29); + TaskId up_spmd_tid_inline264__ssa_v4 = task_29_outs.task_id(); + up_tids_inline310[4] = up_spmd_tid_inline264__ssa_v4; + up_tids_inline310[9] = up_spmd_tid_inline264__ssa_v4; + up_tids_inline310[14] = up_spmd_tid_inline264__ssa_v4; + up_tids_inline310[19] = up_spmd_tid_inline264__ssa_v4; + up_tids_inline310[24] = up_spmd_tid_inline264__ssa_v4; + up_tids_inline310[29] = up_spmd_tid_inline264__ssa_v4; + for (int64_t n_out_inline275 = 6; n_out_inline275 < 17; n_out_inline275 += 1) { + int64_t n0_inline122 = (n_out_inline275 * 1024); + for (int64_t k_split_inline276 = 0; k_split_inline276 < 5; k_split_inline276 += 1) { + int64_t k0_inline113 = (k_split_inline276 * 1024); + TaskId _submit_deps_buf_inline102[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline102[__init_i] = TaskId::invalid(); + TaskId t__tmp_v135 = gate_late_tids_inline249[k_split_inline276]; + _submit_deps_buf_inline102[0] = t__tmp_v135; + + // Task 30: gate_proj_4 + CoreTaskArgs params_t30; + params_t30.add_input(mlp_norm_in_inline71); + params_t30.add_input(ext_w_gate); + // Each split writes its own [16, 1024] partial: rows + // [k_split*16, k_split*16+16) of `gate_part_all`, column band at + // `n0_inline122`. The splits touch disjoint regions, so TensorMap derives + // no edge between them and the kernel stores without an atomic. + const simpler::tmr::Tensor gate_part_all_band_t30 = + gate_part_all + .slice( + 0, static_cast(k_split_inline276 * 16), + static_cast(k_split_inline276 * 16) + 16 + ) + .slice( + 1, static_cast(n0_inline122), + static_cast(n0_inline122) + 1024 + ); + params_t30.add_output(gate_part_all_band_t30); + params_t30.add_scalar(k0_inline113); + params_t30.add_scalar(layer_hidden_base_inline151); + params_t30.add_scalar(n0_inline122); + TaskId params_t30_deps[1]; + uint32_t params_t30_deps_count = 0; + if (_submit_deps_buf_inline102[0].is_valid()) + params_t30_deps[params_t30_deps_count++] = _submit_deps_buf_inline102[0]; + params_t30.set_dependencies(params_t30_deps, params_t30_deps_count); + TaskOutputTensors task_30_outs = rt_submit_aic_task(31, params_t30); + TaskId gate_tid_inline277 = task_30_outs.task_id(); + gate_tids_inline56[((n_out_inline275 * 5) + k_split_inline276)] = gate_tid_inline277; + TaskId _submit_deps_buf_inline246[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline246[__init_i] = TaskId::invalid(); + TaskId t__tmp_v136 = up_late_tids_inline69[k_split_inline276]; + _submit_deps_buf_inline246[0] = t__tmp_v136; + + // Task 31: up_proj_4 + CoreTaskArgs params_t31; + params_t31.add_input(mlp_norm_in_inline71); + params_t31.add_input(ext_w_up); + // Each split writes its own [16, 1024] partial: rows + // [k_split*16, k_split*16+16) of `up_part_all`, column band at + // `n0_inline122`. The splits touch disjoint regions, so TensorMap derives + // no edge between them and the kernel stores without an atomic. + const simpler::tmr::Tensor up_part_all_band_t31 = + up_part_all + .slice( + 0, static_cast(k_split_inline276 * 16), + static_cast(k_split_inline276 * 16) + 16 + ) + .slice( + 1, static_cast(n0_inline122), + static_cast(n0_inline122) + 1024 + ); + params_t31.add_output(up_part_all_band_t31); + params_t31.add_scalar(k0_inline113); + params_t31.add_scalar(layer_hidden_base_inline151); + params_t31.add_scalar(n0_inline122); + TaskId params_t31_deps[1]; + uint32_t params_t31_deps_count = 0; + if (_submit_deps_buf_inline246[0].is_valid()) + params_t31_deps[params_t31_deps_count++] = _submit_deps_buf_inline246[0]; + params_t31.set_dependencies(params_t31_deps, params_t31_deps_count); + TaskOutputTensors task_31_outs = rt_submit_aic_task(32, params_t31); + TaskId up_tid_inline290 = task_31_outs.task_id(); + up_tids_inline310[((n_out_inline275 * 5) + k_split_inline276)] = up_tid_inline290; + } + + // Tasks 30r / 31r: partials_reduce_inter. Bands 6..16 are written only by + // the five split-K tasks above, so each band's sum replaces the accumulator + // rather than joining it. Bands 0..5 come from the SPMD form and are not + // reduced here. + CoreTaskArgs params_t30r; + const simpler::tmr::Tensor gate_part_all_band_t30r = gate_part_all.slice( + 1, static_cast(n0_inline122), static_cast(n0_inline122) + 1024 + ); + params_t30r.add_input(gate_part_all_band_t30r); + const simpler::tmr::Tensor gate_acc_all_inline203_band_t30r = gate_acc_all_inline203.slice( + 1, static_cast(n0_inline122), static_cast(n0_inline122) + 1024 + ); + params_t30r.add_output(gate_acc_all_inline203_band_t30r); + params_t30r.add_scalar(static_cast(5)); + params_t30r.add_scalar(static_cast(1024)); + TaskOutputTensors task_30r_outs = rt_submit_aiv_task(39, params_t30r); + (void)task_30r_outs; + + CoreTaskArgs params_t31r; + const simpler::tmr::Tensor up_part_all_band_t31r = up_part_all.slice( + 1, static_cast(n0_inline122), static_cast(n0_inline122) + 1024 + ); + params_t31r.add_input(up_part_all_band_t31r); + const simpler::tmr::Tensor up_acc_all_inline303_band_t31r = up_acc_all_inline303.slice( + 1, static_cast(n0_inline122), static_cast(n0_inline122) + 1024 + ); + params_t31r.add_output(up_acc_all_inline303_band_t31r); + params_t31r.add_scalar(static_cast(5)); + params_t31r.add_scalar(static_cast(1024)); + TaskOutputTensors task_31r_outs = rt_submit_aiv_task(39, params_t31r); + (void)task_31r_outs; + } + for (int64_t n_out_inline292 = 0; n_out_inline292 < 17; n_out_inline292 += 1) { + int64_t n0_inline122__ssa_v7 = (n_out_inline292 * 1024); + TaskId _submit_deps_buf_inline167[11]; + for (int64_t __init_i = 0; __init_i < 11; ++__init_i) + _submit_deps_buf_inline167[__init_i] = TaskId::invalid(); + _submit_deps_buf_inline167[0] = reduce_tid_inline226; + TaskId t__tmp_v137 = gate_tids_inline56[(n_out_inline292 * 5)]; + _submit_deps_buf_inline167[1] = t__tmp_v137; + TaskId t__tmp_v138 = gate_tids_inline56[((n_out_inline292 * 5) + 1)]; + _submit_deps_buf_inline167[2] = t__tmp_v138; + TaskId t__tmp_v139 = gate_tids_inline56[((n_out_inline292 * 5) + 2)]; + _submit_deps_buf_inline167[3] = t__tmp_v139; + TaskId t__tmp_v140 = gate_tids_inline56[((n_out_inline292 * 5) + 3)]; + _submit_deps_buf_inline167[4] = t__tmp_v140; + TaskId t__tmp_v141 = gate_tids_inline56[((n_out_inline292 * 5) + 4)]; + _submit_deps_buf_inline167[5] = t__tmp_v141; + TaskId t__tmp_v142 = up_tids_inline310[(n_out_inline292 * 5)]; + _submit_deps_buf_inline167[6] = t__tmp_v142; + TaskId t__tmp_v143 = up_tids_inline310[((n_out_inline292 * 5) + 1)]; + _submit_deps_buf_inline167[7] = t__tmp_v143; + TaskId t__tmp_v144 = up_tids_inline310[((n_out_inline292 * 5) + 2)]; + _submit_deps_buf_inline167[8] = t__tmp_v144; + TaskId t__tmp_v145 = up_tids_inline310[((n_out_inline292 * 5) + 3)]; + _submit_deps_buf_inline167[9] = t__tmp_v145; + TaskId t__tmp_v146 = up_tids_inline310[((n_out_inline292 * 5) + 4)]; + _submit_deps_buf_inline167[10] = t__tmp_v146; + + // Task 32: silu + CoreTaskArgs params_t32; + params_t32.add_input(inv_rms_tile_inline126); + // This task writes only the 1024-wide column band at `n0_inline122__ssa_v7`; the + // kernel's store is Shape<...,1024> and its base already carries the + // view's start_offset. Declaring the whole buffer would make TensorMap + // order every band against every other one. + const simpler::tmr::Tensor mlp_tile_inline149_band_t32 = mlp_tile_inline149.slice( + 1, static_cast(n0_inline122__ssa_v7), + static_cast(n0_inline122__ssa_v7) + 1024 + ); + params_t32.add_inout(mlp_tile_inline149_band_t32); + // silu reads only the 1024-wide band at `n0` (same offset it writes to + // mlp_tile); declaring the whole accumulator would depend on every band. + const simpler::tmr::Tensor gate_acc_all_inline203_band_t32 = gate_acc_all_inline203.slice( + 1, static_cast(n0_inline122__ssa_v7), + static_cast(n0_inline122__ssa_v7) + 1024 + ); + params_t32.add_input(gate_acc_all_inline203_band_t32); + // silu reads only the 1024-wide band at `n0` (same offset it writes to + // mlp_tile); declaring the whole accumulator would depend on every band. + const simpler::tmr::Tensor up_acc_all_inline303_band_t32 = up_acc_all_inline303.slice( + 1, static_cast(n0_inline122__ssa_v7), + static_cast(n0_inline122__ssa_v7) + 1024 + ); + params_t32.add_input(up_acc_all_inline303_band_t32); + params_t32.add_scalar(n0_inline122__ssa_v7); + TaskId params_t32_deps[11]; + uint32_t params_t32_deps_count = 0; + if (_submit_deps_buf_inline167[0].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[0]; + if (_submit_deps_buf_inline167[1].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[1]; + if (_submit_deps_buf_inline167[2].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[2]; + if (_submit_deps_buf_inline167[3].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[3]; + if (_submit_deps_buf_inline167[4].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[4]; + if (_submit_deps_buf_inline167[5].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[5]; + if (_submit_deps_buf_inline167[6].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[6]; + if (_submit_deps_buf_inline167[7].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[7]; + if (_submit_deps_buf_inline167[8].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[8]; + if (_submit_deps_buf_inline167[9].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[9]; + if (_submit_deps_buf_inline167[10].is_valid()) + params_t32_deps[params_t32_deps_count++] = _submit_deps_buf_inline167[10]; + params_t32.set_dependencies(params_t32_deps, params_t32_deps_count); + TaskOutputTensors task_32_outs = rt_submit_aiv_task(33, params_t32); + TaskId silu_tid_inline80 = task_32_outs.task_id(); + silu_tids_inline265[n_out_inline292] = silu_tid_inline80; + } + for (int64_t n_out_inline195 = 0; n_out_inline195 < 5; n_out_inline195 += 1) { + int64_t n0_inline122__ssa_v8 = (n_out_inline195 * 1024); + for (int64_t k_split_inline302 = 0; k_split_inline302 < 17; k_split_inline302 += 1) { + int64_t k0_inline113__ssa_v8 = (k_split_inline302 * 1024); + TaskId _submit_deps_buf_inline229[1]; + for (int64_t __init_i = 0; __init_i < 1; ++__init_i) + _submit_deps_buf_inline229[__init_i] = TaskId::invalid(); + TaskId t__tmp_v152 = silu_tids_inline265[k_split_inline302]; + _submit_deps_buf_inline229[0] = t__tmp_v152; + + // Task 33: down_proj + CoreTaskArgs params_t33; + // down_proj reads only the 1024-wide band at `k0` (its 16 x 64 loads span + // k0 .. k0+1024), which is exactly the band one silu task writes. + const simpler::tmr::Tensor mlp_tile_inline149_band_t33 = mlp_tile_inline149.slice( + 1, static_cast(k0_inline113__ssa_v8), + static_cast(k0_inline113__ssa_v8) + 1024 + ); + params_t33.add_input(mlp_tile_inline149_band_t33); + params_t33.add_input(ext_w_down); + // Each split writes its own [16, 1024] partial: rows + // [k_split*16, k_split*16+16) of `down_part_all`, column band at + // `n0_inline122__ssa_v8`. The splits touch disjoint regions, so TensorMap + // derives no edge between them and the kernel stores without an atomic. + const simpler::tmr::Tensor down_part_all_band_t33 = + down_part_all + .slice( + 0, static_cast(k_split_inline302 * 16), + static_cast(k_split_inline302 * 16) + 16 + ) + .slice( + 1, static_cast(n0_inline122__ssa_v8), + static_cast(n0_inline122__ssa_v8) + 1024 + ); + params_t33.add_output(down_part_all_band_t33); + params_t33.add_scalar(k0_inline113__ssa_v8); + params_t33.add_scalar(layer_inter_base_inline107); + params_t33.add_scalar(n0_inline122__ssa_v8); + TaskId params_t33_deps[1]; + uint32_t params_t33_deps_count = 0; + if (_submit_deps_buf_inline229[0].is_valid()) + params_t33_deps[params_t33_deps_count++] = _submit_deps_buf_inline229[0]; + params_t33.set_dependencies(params_t33_deps, params_t33_deps_count); + params_t33.set_allow_early_resolve(true); + TaskOutputTensors task_33_outs = rt_submit_aic_task(34, params_t33); + TaskId down_tid_inline210 = task_33_outs.task_id(); + down_tids_inline156[((n_out_inline195 * 17) + k_split_inline302)] = down_tid_inline210; + } + + // Task 33r: down_partials_reduce. Sums this band's 17 private partials into + // the accumulator. The input view spans all 17 partials' rows, so TensorMap + // gives this task 17 unordered producers: the band's readiness chain is one + // parallel round plus this task, where the atomic form was 17 deep. + CoreTaskArgs params_t33r; + const simpler::tmr::Tensor down_part_all_band_t33r = down_part_all.slice( + 1, static_cast(n0_inline122__ssa_v8), + static_cast(n0_inline122__ssa_v8) + 1024 + ); + params_t33r.add_input(down_part_all_band_t33r); + const simpler::tmr::Tensor down_acc_all_inline168_band_t33r = down_acc_all_inline168.slice( + 1, static_cast(n0_inline122__ssa_v8), + static_cast(n0_inline122__ssa_v8) + 1024 + ); + params_t33r.add_output(down_acc_all_inline168_band_t33r); + params_t33r.add_scalar(static_cast(17)); + params_t33r.add_scalar(static_cast(1024)); + TaskOutputTensors task_33r_outs = rt_submit_aiv_task(37, params_t33r); + (void)task_33r_outs; + } + } + TaskId _submit_deps_buf_inline123[85]; + for (int64_t __init_i = 0; __init_i < 85; ++__init_i) + _submit_deps_buf_inline123[__init_i] = TaskId::invalid(); + TaskId t__tmp_v153 = down_tids_inline156[0]; + _submit_deps_buf_inline123[0] = t__tmp_v153; + TaskId t__tmp_v154 = down_tids_inline156[1]; + _submit_deps_buf_inline123[1] = t__tmp_v154; + TaskId t__tmp_v155 = down_tids_inline156[2]; + _submit_deps_buf_inline123[2] = t__tmp_v155; + TaskId t__tmp_v156 = down_tids_inline156[3]; + _submit_deps_buf_inline123[3] = t__tmp_v156; + TaskId t__tmp_v157 = down_tids_inline156[4]; + _submit_deps_buf_inline123[4] = t__tmp_v157; + TaskId t__tmp_v158 = down_tids_inline156[5]; + _submit_deps_buf_inline123[5] = t__tmp_v158; + TaskId t__tmp_v159 = down_tids_inline156[6]; + _submit_deps_buf_inline123[6] = t__tmp_v159; + TaskId t__tmp_v160 = down_tids_inline156[7]; + _submit_deps_buf_inline123[7] = t__tmp_v160; + TaskId t__tmp_v161 = down_tids_inline156[8]; + _submit_deps_buf_inline123[8] = t__tmp_v161; + TaskId t__tmp_v162 = down_tids_inline156[9]; + _submit_deps_buf_inline123[9] = t__tmp_v162; + TaskId t__tmp_v163 = down_tids_inline156[10]; + _submit_deps_buf_inline123[10] = t__tmp_v163; + TaskId t__tmp_v164 = down_tids_inline156[11]; + _submit_deps_buf_inline123[11] = t__tmp_v164; + TaskId t__tmp_v165 = down_tids_inline156[12]; + _submit_deps_buf_inline123[12] = t__tmp_v165; + TaskId t__tmp_v166 = down_tids_inline156[13]; + _submit_deps_buf_inline123[13] = t__tmp_v166; + TaskId t__tmp_v167 = down_tids_inline156[14]; + _submit_deps_buf_inline123[14] = t__tmp_v167; + TaskId t__tmp_v168 = down_tids_inline156[15]; + _submit_deps_buf_inline123[15] = t__tmp_v168; + TaskId t__tmp_v169 = down_tids_inline156[16]; + _submit_deps_buf_inline123[16] = t__tmp_v169; + TaskId t__tmp_v170 = down_tids_inline156[17]; + _submit_deps_buf_inline123[17] = t__tmp_v170; + TaskId t__tmp_v171 = down_tids_inline156[18]; + _submit_deps_buf_inline123[18] = t__tmp_v171; + TaskId t__tmp_v172 = down_tids_inline156[19]; + _submit_deps_buf_inline123[19] = t__tmp_v172; + TaskId t__tmp_v173 = down_tids_inline156[20]; + _submit_deps_buf_inline123[20] = t__tmp_v173; + TaskId t__tmp_v174 = down_tids_inline156[21]; + _submit_deps_buf_inline123[21] = t__tmp_v174; + TaskId t__tmp_v175 = down_tids_inline156[22]; + _submit_deps_buf_inline123[22] = t__tmp_v175; + TaskId t__tmp_v176 = down_tids_inline156[23]; + _submit_deps_buf_inline123[23] = t__tmp_v176; + TaskId t__tmp_v177 = down_tids_inline156[24]; + _submit_deps_buf_inline123[24] = t__tmp_v177; + TaskId t__tmp_v178 = down_tids_inline156[25]; + _submit_deps_buf_inline123[25] = t__tmp_v178; + TaskId t__tmp_v179 = down_tids_inline156[26]; + _submit_deps_buf_inline123[26] = t__tmp_v179; + TaskId t__tmp_v180 = down_tids_inline156[27]; + _submit_deps_buf_inline123[27] = t__tmp_v180; + TaskId t__tmp_v181 = down_tids_inline156[28]; + _submit_deps_buf_inline123[28] = t__tmp_v181; + TaskId t__tmp_v182 = down_tids_inline156[29]; + _submit_deps_buf_inline123[29] = t__tmp_v182; + TaskId t__tmp_v183 = down_tids_inline156[30]; + _submit_deps_buf_inline123[30] = t__tmp_v183; + TaskId t__tmp_v184 = down_tids_inline156[31]; + _submit_deps_buf_inline123[31] = t__tmp_v184; + TaskId t__tmp_v185 = down_tids_inline156[32]; + _submit_deps_buf_inline123[32] = t__tmp_v185; + TaskId t__tmp_v186 = down_tids_inline156[33]; + _submit_deps_buf_inline123[33] = t__tmp_v186; + TaskId t__tmp_v187 = down_tids_inline156[34]; + _submit_deps_buf_inline123[34] = t__tmp_v187; + TaskId t__tmp_v188 = down_tids_inline156[35]; + _submit_deps_buf_inline123[35] = t__tmp_v188; + TaskId t__tmp_v189 = down_tids_inline156[36]; + _submit_deps_buf_inline123[36] = t__tmp_v189; + TaskId t__tmp_v190 = down_tids_inline156[37]; + _submit_deps_buf_inline123[37] = t__tmp_v190; + TaskId t__tmp_v191 = down_tids_inline156[38]; + _submit_deps_buf_inline123[38] = t__tmp_v191; + TaskId t__tmp_v192 = down_tids_inline156[39]; + _submit_deps_buf_inline123[39] = t__tmp_v192; + TaskId t__tmp_v193 = down_tids_inline156[40]; + _submit_deps_buf_inline123[40] = t__tmp_v193; + TaskId t__tmp_v194 = down_tids_inline156[41]; + _submit_deps_buf_inline123[41] = t__tmp_v194; + TaskId t__tmp_v195 = down_tids_inline156[42]; + _submit_deps_buf_inline123[42] = t__tmp_v195; + TaskId t__tmp_v196 = down_tids_inline156[43]; + _submit_deps_buf_inline123[43] = t__tmp_v196; + TaskId t__tmp_v197 = down_tids_inline156[44]; + _submit_deps_buf_inline123[44] = t__tmp_v197; + TaskId t__tmp_v198 = down_tids_inline156[45]; + _submit_deps_buf_inline123[45] = t__tmp_v198; + TaskId t__tmp_v199 = down_tids_inline156[46]; + _submit_deps_buf_inline123[46] = t__tmp_v199; + TaskId t__tmp_v200 = down_tids_inline156[47]; + _submit_deps_buf_inline123[47] = t__tmp_v200; + TaskId t__tmp_v201 = down_tids_inline156[48]; + _submit_deps_buf_inline123[48] = t__tmp_v201; + TaskId t__tmp_v202 = down_tids_inline156[49]; + _submit_deps_buf_inline123[49] = t__tmp_v202; + TaskId t__tmp_v203 = down_tids_inline156[50]; + _submit_deps_buf_inline123[50] = t__tmp_v203; + TaskId t__tmp_v204 = down_tids_inline156[51]; + _submit_deps_buf_inline123[51] = t__tmp_v204; + TaskId t__tmp_v205 = down_tids_inline156[52]; + _submit_deps_buf_inline123[52] = t__tmp_v205; + TaskId t__tmp_v206 = down_tids_inline156[53]; + _submit_deps_buf_inline123[53] = t__tmp_v206; + TaskId t__tmp_v207 = down_tids_inline156[54]; + _submit_deps_buf_inline123[54] = t__tmp_v207; + TaskId t__tmp_v208 = down_tids_inline156[55]; + _submit_deps_buf_inline123[55] = t__tmp_v208; + TaskId t__tmp_v209 = down_tids_inline156[56]; + _submit_deps_buf_inline123[56] = t__tmp_v209; + TaskId t__tmp_v210 = down_tids_inline156[57]; + _submit_deps_buf_inline123[57] = t__tmp_v210; + TaskId t__tmp_v211 = down_tids_inline156[58]; + _submit_deps_buf_inline123[58] = t__tmp_v211; + TaskId t__tmp_v212 = down_tids_inline156[59]; + _submit_deps_buf_inline123[59] = t__tmp_v212; + TaskId t__tmp_v213 = down_tids_inline156[60]; + _submit_deps_buf_inline123[60] = t__tmp_v213; + TaskId t__tmp_v214 = down_tids_inline156[61]; + _submit_deps_buf_inline123[61] = t__tmp_v214; + TaskId t__tmp_v215 = down_tids_inline156[62]; + _submit_deps_buf_inline123[62] = t__tmp_v215; + TaskId t__tmp_v216 = down_tids_inline156[63]; + _submit_deps_buf_inline123[63] = t__tmp_v216; + TaskId t__tmp_v217 = down_tids_inline156[64]; + _submit_deps_buf_inline123[64] = t__tmp_v217; + TaskId t__tmp_v218 = down_tids_inline156[65]; + _submit_deps_buf_inline123[65] = t__tmp_v218; + TaskId t__tmp_v219 = down_tids_inline156[66]; + _submit_deps_buf_inline123[66] = t__tmp_v219; + TaskId t__tmp_v220 = down_tids_inline156[67]; + _submit_deps_buf_inline123[67] = t__tmp_v220; + TaskId t__tmp_v221 = down_tids_inline156[68]; + _submit_deps_buf_inline123[68] = t__tmp_v221; + TaskId t__tmp_v222 = down_tids_inline156[69]; + _submit_deps_buf_inline123[69] = t__tmp_v222; + TaskId t__tmp_v223 = down_tids_inline156[70]; + _submit_deps_buf_inline123[70] = t__tmp_v223; + TaskId t__tmp_v224 = down_tids_inline156[71]; + _submit_deps_buf_inline123[71] = t__tmp_v224; + TaskId t__tmp_v225 = down_tids_inline156[72]; + _submit_deps_buf_inline123[72] = t__tmp_v225; + TaskId t__tmp_v226 = down_tids_inline156[73]; + _submit_deps_buf_inline123[73] = t__tmp_v226; + TaskId t__tmp_v227 = down_tids_inline156[74]; + _submit_deps_buf_inline123[74] = t__tmp_v227; + TaskId t__tmp_v228 = down_tids_inline156[75]; + _submit_deps_buf_inline123[75] = t__tmp_v228; + TaskId t__tmp_v229 = down_tids_inline156[76]; + _submit_deps_buf_inline123[76] = t__tmp_v229; + TaskId t__tmp_v230 = down_tids_inline156[77]; + _submit_deps_buf_inline123[77] = t__tmp_v230; + TaskId t__tmp_v231 = down_tids_inline156[78]; + _submit_deps_buf_inline123[78] = t__tmp_v231; + TaskId t__tmp_v232 = down_tids_inline156[79]; + _submit_deps_buf_inline123[79] = t__tmp_v232; + TaskId t__tmp_v233 = down_tids_inline156[80]; + _submit_deps_buf_inline123[80] = t__tmp_v233; + TaskId t__tmp_v234 = down_tids_inline156[81]; + _submit_deps_buf_inline123[81] = t__tmp_v234; + TaskId t__tmp_v235 = down_tids_inline156[82]; + _submit_deps_buf_inline123[82] = t__tmp_v235; + TaskId t__tmp_v236 = down_tids_inline156[83]; + _submit_deps_buf_inline123[83] = t__tmp_v236; + TaskId t__tmp_v237 = down_tids_inline156[84]; + _submit_deps_buf_inline123[84] = t__tmp_v237; + + // Spmd dcr_xgamma_spmd: dcr_xgamma + CoreTaskArgs params_t34; + params_t34.add_input(down_acc_all_inline168); + params_t34.add_input(post_norm_partial_inline118); + params_t34.add_inout(next_hidden); + params_t34.add_input(ext_input_rms_weight); + params_t34.add_inout(next_normed); + params_t34.add_scalar(next_gamma_idx); + params_t34.launch_spec.set_block_num(5); + params_t34.set_allow_early_resolve(true); + TaskId params_t34_deps[85]; + uint32_t params_t34_deps_count = 0; + if (_submit_deps_buf_inline123[0].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[0]; + if (_submit_deps_buf_inline123[1].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[1]; + if (_submit_deps_buf_inline123[2].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[2]; + if (_submit_deps_buf_inline123[3].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[3]; + if (_submit_deps_buf_inline123[4].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[4]; + if (_submit_deps_buf_inline123[5].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[5]; + if (_submit_deps_buf_inline123[6].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[6]; + if (_submit_deps_buf_inline123[7].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[7]; + if (_submit_deps_buf_inline123[8].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[8]; + if (_submit_deps_buf_inline123[9].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[9]; + if (_submit_deps_buf_inline123[10].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[10]; + if (_submit_deps_buf_inline123[11].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[11]; + if (_submit_deps_buf_inline123[12].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[12]; + if (_submit_deps_buf_inline123[13].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[13]; + if (_submit_deps_buf_inline123[14].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[14]; + if (_submit_deps_buf_inline123[15].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[15]; + if (_submit_deps_buf_inline123[16].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[16]; + if (_submit_deps_buf_inline123[17].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[17]; + if (_submit_deps_buf_inline123[18].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[18]; + if (_submit_deps_buf_inline123[19].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[19]; + if (_submit_deps_buf_inline123[20].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[20]; + if (_submit_deps_buf_inline123[21].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[21]; + if (_submit_deps_buf_inline123[22].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[22]; + if (_submit_deps_buf_inline123[23].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[23]; + if (_submit_deps_buf_inline123[24].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[24]; + if (_submit_deps_buf_inline123[25].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[25]; + if (_submit_deps_buf_inline123[26].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[26]; + if (_submit_deps_buf_inline123[27].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[27]; + if (_submit_deps_buf_inline123[28].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[28]; + if (_submit_deps_buf_inline123[29].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[29]; + if (_submit_deps_buf_inline123[30].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[30]; + if (_submit_deps_buf_inline123[31].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[31]; + if (_submit_deps_buf_inline123[32].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[32]; + if (_submit_deps_buf_inline123[33].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[33]; + if (_submit_deps_buf_inline123[34].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[34]; + if (_submit_deps_buf_inline123[35].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[35]; + if (_submit_deps_buf_inline123[36].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[36]; + if (_submit_deps_buf_inline123[37].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[37]; + if (_submit_deps_buf_inline123[38].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[38]; + if (_submit_deps_buf_inline123[39].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[39]; + if (_submit_deps_buf_inline123[40].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[40]; + if (_submit_deps_buf_inline123[41].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[41]; + if (_submit_deps_buf_inline123[42].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[42]; + if (_submit_deps_buf_inline123[43].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[43]; + if (_submit_deps_buf_inline123[44].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[44]; + if (_submit_deps_buf_inline123[45].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[45]; + if (_submit_deps_buf_inline123[46].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[46]; + if (_submit_deps_buf_inline123[47].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[47]; + if (_submit_deps_buf_inline123[48].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[48]; + if (_submit_deps_buf_inline123[49].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[49]; + if (_submit_deps_buf_inline123[50].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[50]; + if (_submit_deps_buf_inline123[51].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[51]; + if (_submit_deps_buf_inline123[52].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[52]; + if (_submit_deps_buf_inline123[53].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[53]; + if (_submit_deps_buf_inline123[54].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[54]; + if (_submit_deps_buf_inline123[55].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[55]; + if (_submit_deps_buf_inline123[56].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[56]; + if (_submit_deps_buf_inline123[57].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[57]; + if (_submit_deps_buf_inline123[58].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[58]; + if (_submit_deps_buf_inline123[59].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[59]; + if (_submit_deps_buf_inline123[60].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[60]; + if (_submit_deps_buf_inline123[61].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[61]; + if (_submit_deps_buf_inline123[62].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[62]; + if (_submit_deps_buf_inline123[63].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[63]; + if (_submit_deps_buf_inline123[64].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[64]; + if (_submit_deps_buf_inline123[65].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[65]; + if (_submit_deps_buf_inline123[66].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[66]; + if (_submit_deps_buf_inline123[67].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[67]; + if (_submit_deps_buf_inline123[68].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[68]; + if (_submit_deps_buf_inline123[69].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[69]; + if (_submit_deps_buf_inline123[70].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[70]; + if (_submit_deps_buf_inline123[71].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[71]; + if (_submit_deps_buf_inline123[72].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[72]; + if (_submit_deps_buf_inline123[73].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[73]; + if (_submit_deps_buf_inline123[74].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[74]; + if (_submit_deps_buf_inline123[75].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[75]; + if (_submit_deps_buf_inline123[76].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[76]; + if (_submit_deps_buf_inline123[77].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[77]; + if (_submit_deps_buf_inline123[78].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[78]; + if (_submit_deps_buf_inline123[79].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[79]; + if (_submit_deps_buf_inline123[80].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[80]; + if (_submit_deps_buf_inline123[81].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[81]; + if (_submit_deps_buf_inline123[82].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[82]; + if (_submit_deps_buf_inline123[83].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[83]; + if (_submit_deps_buf_inline123[84].is_valid()) + params_t34_deps[params_t34_deps_count++] = _submit_deps_buf_inline123[84]; + params_t34.set_dependencies(params_t34_deps, params_t34_deps_count); + TaskOutputTensors task_34_outs = rt_submit_aiv_task(35, params_t34); + TaskId dcr_tid_inline58 = task_34_outs.task_id(); + prev_out_tid[0] = dcr_tid_inline58; + prev_normed_tid[0] = dcr_tid_inline58; + simpler::tmr::Tensor cur__ssa_v8 = next_hidden; + simpler::tmr::Tensor normed__ssa_v6 = next_normed; + cur__rv_v7 = cur__ssa_v8; + normed__rv_v5 = normed__ssa_v6; + } + } + for (int64_t ob0 = 0; ob0 < 16; ob0 += 16) { + SIMPLER_SCOPE() { + // Task 35: copy_out + CoreTaskArgs params_t35; + params_t35.add_output(ext_out); + params_t35.add_input(cur__rv_v7); + params_t35.add_scalar(ob0); + rt_submit_aiv_task(36, params_t35); + } + } + } +} + +} // extern "C" diff --git a/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/test_qwen3_14b_decode_auto.py b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/test_qwen3_14b_decode_auto.py new file mode 100644 index 0000000000..b52d20cd12 --- /dev/null +++ b/examples/a2a3/tensormap_and_ringbuffer/qwen3_14b_decode_auto/test_qwen3_14b_decode_auto.py @@ -0,0 +1,470 @@ +#!/usr/bin/env python3 +# Copyright (c) PyPTO Contributors. +# This program is free software, you can redistribute it and/or modify it under the terms and conditions of +# CANN Open Software License Agreement Version 2.0 (the "License"). +# Please refer to the License for details. You may not use this file except in compliance with the License. +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. +# See LICENSE in the root of the software repository for the full text of the License. +# ----------------------------------------------------------------------------------------------------------- +"""Qwen3-14B 40-layer decode under TensorMap-derived dependencies — SceneTestCase. + +The AUTO-scope twin of ``../qwen3_14b_decode/``: same network, same fixture, +same golden, but ``SIMPLER_SCOPE(ScopeMode::AUTO)`` instead of MANUAL, so every +WAIT edge is derived from tensor overlap rather than declared by hand. Its +``INOUT`` / ``INPUT`` views are ``.slice()``-narrowed to the column band each +task indexes, because a whole-buffer declaration under AUTO serializes bands +that do not overlap. + +Four split-K accumulators write private partials that a ``partials_reduce_*`` +task sums, rather than each atomically adding into a shared band: the atomic form +is commutative but ``INOUT`` cannot say so, so TensorMap had to serialize the +writers. + +Only the 13 incores that had to change live here; the ``CALLABLE`` reaches the +other 27 as ``../qwen3_14b_decode/kernels/...`` so the harvested codegen has one +home. See README.md for the measured A/B and what the residual critical path is +made of; see the sibling's README for model, provenance and regen. + +Parameter regime matches ``stress_profile.py`` (vLLM serving stress): BATCH=16, +MAX_SEQ=5500 (= max_model_len), fixed decode seq_len=3500. Weights and the paged +KV pool are stacked x40 (one slice per layer); every layer reuses layer-0 +weights, per the lib const-layer-0 stacked-fwd reference, while each layer reads +and writes its own KV pool. +""" + +from simpler.task_interface import ArgDirection as D + +from simpler_setup import SceneTestCase, scene_test +from simpler_setup.goldens.qwen3_14b_decode import ( + compute_golden as _decode_golden, +) +from simpler_setup.goldens.qwen3_14b_decode import ( + generate_inputs as _decode_generate_inputs, +) + +# CANN devkit headers for the attention extern, which builds on AscendC and the +# vendored FusedInferAttentionScore under kernels/paged_attention_cce/vendor/. +# `vendor/.../attn_infra/base_defs.hpp` selects its AscendC entry header under +# `#if ASC_DEVKIT_MAJOR >= 9`, which ccec predefines from the installed devkit, +# so a CANN 9 box must be able to resolve `basic_api/kernel_basic_intf.h` from +# one of these. +# +# `$ASCEND_HOME_PATH` keeps the paths machine-independent and is expanded at +# compile time, not import time — this file is collected on sim and macOS runners +# that have no CANN at all. The devkit's arch subdirectory is named differently +# across installs, so both layouts are listed; missing ones are dropped, and the +# scene-test resolver raises if *every* entry is missing. +_CANN_SUBDIRS = ( + "include", + "asc", + "asc/impl/adv_api", + "asc/impl/basic_api", + "asc/impl/basic_api/reg_compute", + "asc/impl/c_api", + "asc/impl/simt_api", + "asc/impl/utils", + "asc/include", + "asc/include/adv_api", + "asc/include/aicpu_api", + "asc/include/basic_api", + "asc/include/basic_api/reg_compute", + "asc/include/c_api", + "asc/include/interface", + "asc/include/simt_api", + "asc/include/utils", + "tikcpp/tikcfw", + "tikcpp/tikcfw/impl", + "tikcpp/tikcfw/interface", +) + +_CANN_INCLUDE_DIRS = [f"$ASCEND_HOME_PATH/{prefix}{sub}" for prefix in ("aarch64-linux/", "") for sub in _CANN_SUBDIRS] + + +# Validates the full 40-layer fused decode against a torch reference. +@scene_test(level=2, runtime="tensormap_and_ringbuffer") +class TestQwen314BDecodeAuto(SceneTestCase): + """Qwen3-14B decode, all 40 layers in one dispatch, against a torch reference.""" + + RTOL = 5e-2 + ATOL = 1e-1 + + CALLABLE = { + "orchestration": { + "source": "kernels/orchestration/decode_fwd_layers.cpp", + "function_name": "aicpu_orchestration_entry", + # decode_fwd_layers takes k_cache / v_cache as plain inputs, but the + # attention extern writes the current token's KV into them. Declaring + # them INOUT here has simpler copy the pools back, so the golden can + # check all 40 layers' KV writes and not just the hidden output. + "signature": [ + D.IN, # 0 hidden_states + D.IN, # 1 input_rms_weight + D.IN, # 2 wq + D.IN, # 3 wk + D.IN, # 4 wv + D.IN, # 5 q_norm_weight + D.IN, # 6 k_norm_weight + D.IN, # 7 seq_lens + D.IN, # 8 block_table + D.IN, # 9 slot_mapping + D.IN, # 10 rope_cos + D.IN, # 11 rope_sin + D.INOUT, # 12 k_cache + D.INOUT, # 13 v_cache + D.IN, # 14 wo + D.IN, # 15 w_gate + D.IN, # 16 w_up + D.IN, # 17 w_down + D.IN, # 18 post_rms_weight + D.OUT, # 19 out + ], + }, + # func_id 0..36 are transcribed from the pypto codegen kernel_config.py + # for decode_fwd_layers (N=40); 0/11/12 are the CANN attention externs, and + # 11 and 12 are the same source dispatched as the AIC and AIV halves of one + # mixed task. func_id 37..39 are this example's own split-K partial reducers + # and have no counterpart in the codegen. 13/31/32/34 keep their codegen + # func_id but store a private partial, so their last argument is OUT. + "incores": [ + { + "func_id": 0, + "name": "paged_attention_tiling_cce", + "source": "../qwen3_14b_decode/kernels/vendor/paged_attention_cce/tiling/entry.cpp", + "core_type": "aiv", + "extra_include_dirs": _CANN_INCLUDE_DIRS, + "signature": [D.IN, D.OUT], + }, + { + "func_id": 1, + "name": "copy_hidden", + "source": "../qwen3_14b_decode/kernels/aiv/copy_hidden.cpp", + "core_type": "aiv", + "signature": [D.OUT, D.IN], + }, + { + "func_id": 2, + "name": "x_gamma0", + "source": "../qwen3_14b_decode/kernels/aiv/x_gamma0.cpp", + "core_type": "aiv", + "signature": [D.OUT, D.IN, D.IN], + }, + { + "func_id": 3, + "name": "attn_out_seed", + "source": "../qwen3_14b_decode/kernels/aiv/attn_out_seed.cpp", + "core_type": "aiv", + "signature": [D.IN], + }, + { + "func_id": 4, + "name": "rms_recip", + "source": "../qwen3_14b_decode/kernels/aiv/rms_recip.cpp", + "core_type": "aiv", + "signature": [D.IN, D.INOUT], + }, + { + "func_id": 5, + "name": "q_seed", + "source": "../qwen3_14b_decode/kernels/aiv/q_seed.cpp", + "core_type": "aiv", + "signature": [D.INOUT], + }, + { + "func_id": 6, + "name": "q_proj", + "source": "../qwen3_14b_decode/kernels/aic/q_proj.cpp", + "core_type": "aic", + "signature": [D.INOUT, D.IN, D.IN], + }, + { + "func_id": 7, + "name": "kv_seed", + "source": "../qwen3_14b_decode/kernels/aiv/kv_seed.cpp", + "core_type": "aiv", + "signature": [D.INOUT, D.INOUT], + }, + { + "func_id": 8, + "name": "mlp_out_seed", + "source": "../qwen3_14b_decode/kernels/aiv/mlp_out_seed.cpp", + "core_type": "aiv", + "signature": [D.INOUT, D.INOUT, D.INOUT, D.INOUT], + }, + { + "func_id": 9, + "name": "k_proj", + "source": "../qwen3_14b_decode/kernels/aic/k_proj.cpp", + "core_type": "aic", + "signature": [D.INOUT, D.IN, D.IN], + }, + { + "func_id": 10, + "name": "v_proj", + "source": "../qwen3_14b_decode/kernels/aic/v_proj.cpp", + "core_type": "aic", + "signature": [D.INOUT, D.IN, D.IN], + }, + { + "func_id": 11, + "name": "paged_attention_rope_cce_aic", + "source": "../qwen3_14b_decode/kernels/vendor/paged_attention_cce/attention_rope/entry.cpp", + "core_type": "aic", + "extra_include_dirs": _CANN_INCLUDE_DIRS, + "signature": [ + D.INOUT, + D.INOUT, + D.INOUT, + D.INOUT, + D.IN, + D.INOUT, + D.INOUT, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + ], + }, + { + "func_id": 12, + "name": "paged_attention_rope_cce_aiv", + "source": "../qwen3_14b_decode/kernels/vendor/paged_attention_cce/attention_rope/entry.cpp", + "core_type": "aiv", + "extra_include_dirs": _CANN_INCLUDE_DIRS, + "signature": [ + D.INOUT, + D.INOUT, + D.INOUT, + D.INOUT, + D.IN, + D.INOUT, + D.INOUT, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + D.IN, + ], + }, + { + "func_id": 13, + "name": "out_proj", + "source": "kernels/aic/out_proj.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.OUT], + }, + { + "func_id": 14, + "name": "out_proj_0", + "source": "../qwen3_14b_decode/kernels/aic/out_proj_0.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 15, + "name": "residual_rms_cast", + "source": "kernels/aiv/residual_rms_cast.cpp", + "core_type": "aiv", + "signature": [D.INOUT, D.INOUT, D.IN, D.IN, D.IN], + }, + { + "func_id": 16, + "name": "residual_rms_cast_0", + "source": "kernels/aiv/residual_rms_cast_0.cpp", + "core_type": "aiv", + "signature": [D.INOUT, D.INOUT, D.IN, D.IN, D.IN], + }, + { + "func_id": 17, + "name": "residual_rms_cast_1", + "source": "kernels/aiv/residual_rms_cast_1.cpp", + "core_type": "aiv", + "signature": [D.INOUT, D.INOUT, D.IN, D.IN, D.IN], + }, + { + "func_id": 18, + "name": "residual_rms_cast_2", + "source": "kernels/aiv/residual_rms_cast_2.cpp", + "core_type": "aiv", + "signature": [D.INOUT, D.INOUT, D.IN, D.IN, D.IN], + }, + { + "func_id": 19, + "name": "residual_rms_cast_3", + "source": "kernels/aiv/residual_rms_cast_3.cpp", + "core_type": "aiv", + "signature": [D.INOUT, D.INOUT, D.IN, D.IN, D.IN], + }, + { + "func_id": 20, + "name": "post_rms_reduce", + "source": "../qwen3_14b_decode/kernels/aiv/post_rms_reduce.cpp", + "core_type": "aiv", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 21, + "name": "gate_proj", + "source": "../qwen3_14b_decode/kernels/aic/gate_proj.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 22, + "name": "up_proj", + "source": "../qwen3_14b_decode/kernels/aic/up_proj.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 23, + "name": "gate_proj_0", + "source": "../qwen3_14b_decode/kernels/aic/gate_proj_0.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 24, + "name": "up_proj_0", + "source": "../qwen3_14b_decode/kernels/aic/up_proj_0.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 25, + "name": "gate_proj_1", + "source": "../qwen3_14b_decode/kernels/aic/gate_proj_1.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 26, + "name": "up_proj_1", + "source": "../qwen3_14b_decode/kernels/aic/up_proj_1.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 27, + "name": "gate_proj_2", + "source": "../qwen3_14b_decode/kernels/aic/gate_proj_2.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 28, + "name": "up_proj_2", + "source": "../qwen3_14b_decode/kernels/aic/up_proj_2.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 29, + "name": "gate_proj_3", + "source": "../qwen3_14b_decode/kernels/aic/gate_proj_3.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 30, + "name": "up_proj_3", + "source": "../qwen3_14b_decode/kernels/aic/up_proj_3.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.INOUT], + }, + { + "func_id": 31, + "name": "gate_proj_4", + "source": "kernels/aic/gate_proj_4.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.OUT], + }, + { + "func_id": 32, + "name": "up_proj_4", + "source": "kernels/aic/up_proj_4.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.OUT], + }, + { + "func_id": 33, + "name": "silu", + "source": "kernels/aiv/silu.cpp", + "core_type": "aiv", + "signature": [D.IN, D.INOUT, D.IN, D.IN], + }, + { + "func_id": 34, + "name": "down_proj", + "source": "kernels/aic/down_proj.cpp", + "core_type": "aic", + "signature": [D.IN, D.IN, D.OUT], + }, + { + "func_id": 37, + "name": "partials_reduce_hidden", + "source": "kernels/aiv/partials_reduce_hidden.cpp", + "core_type": "aiv", + "signature": [D.IN, D.OUT], + }, + { + "func_id": 38, + "name": "partials_reduce_hidden_add", + "source": "kernels/aiv/partials_reduce_hidden_add.cpp", + "core_type": "aiv", + "signature": [D.IN, D.INOUT], + }, + { + "func_id": 39, + "name": "partials_reduce_inter", + "source": "kernels/aiv/partials_reduce_inter.cpp", + "core_type": "aiv", + "signature": [D.IN, D.OUT], + }, + { + "func_id": 35, + "name": "dcr_xgamma", + "source": "../qwen3_14b_decode/kernels/aiv/dcr_xgamma.cpp", + "core_type": "aiv", + "signature": [D.IN, D.IN, D.INOUT, D.IN, D.INOUT], + }, + { + "func_id": 36, + "name": "copy_out", + "source": "../qwen3_14b_decode/kernels/aiv/copy_out.cpp", + "core_type": "aiv", + "signature": [D.OUT, D.IN], + }, + ], + } + + CASES = [ + { + "name": "AutoDepBatch16Seq3500", + "platforms": ["a2a3"], + "manual": True, + # A run takes the whole device, matching the lib default. + "params": {"seed": 1234, "seq_len": 3500}, + }, + ] + + def generate_args(self, params): + return _decode_generate_inputs(params.get("seed", 1234), params.get("seq_len", 3500)) + + def compute_golden(self, args, params): + _decode_golden(args) + + +if __name__ == "__main__": + SceneTestCase.run_module(__name__)