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__)