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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
/*
* 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.
* -----------------------------------------------------------------------------------------------------------
*/

#include <cstdint>

#ifndef __gm__
#define __gm__
#endif
#ifndef __aicore__
#define __aicore__ [aicore]
#endif

#include <pto/pto-inst.hpp>

#include "tensor.h"

using namespace pto;

extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) {
__gm__ Tensor *src_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]);
__gm__ Tensor *result_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]);

__gm__ float *src = reinterpret_cast<__gm__ float *>(src_tensor->buffer.addr) + src_tensor->start_offset;
__gm__ float *result = reinterpret_cast<__gm__ float *>(result_tensor->buffer.addr) + result_tensor->start_offset;

constexpr int kTotalRows = 128;
constexpr int kRows = 64;
constexpr int kCols = 128;
constexpr int kIters = kTotalRows / kRows;
using DynShapeDim5 = Shape<1, 1, 1, kRows, kCols>;
using DynStrideDim5 = pto::Stride<1, 1, 1, kCols, 1>;
using GlobalData = GlobalTensor<float, DynShapeDim5, DynStrideDim5>;
using TileData = Tile<TileType::Vec, float, kRows, kCols, BLayout::RowMajor, -1, -1>;

TileData src_tile(kRows, kCols);
TileData result_tile(kRows, kCols);
TASSIGN(src_tile, 0x0);
TASSIGN(result_tile, 0x10000);

constexpr int kChunkElems = kRows * kCols;
for (int iter = 0; iter < kIters; ++iter) {
GlobalData src_global(src + iter * kChunkElems);
GlobalData result_global(result + iter * kChunkElems);
TLOAD(src_tile, src_global);
set_flag(PIPE_MTE2, PIPE_V, EVENT_ID0);
wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0);

TADDS(result_tile, src_tile, 1.0f);
set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);

TSTORE(result_global, result_tile);
set_flag(PIPE_MTE3, PIPE_S, EVENT_ID7);
wait_flag(PIPE_MTE3, PIPE_S, EVENT_ID7);
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
/*
* 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.
* -----------------------------------------------------------------------------------------------------------
*/

#include <cstdint>

#ifndef __gm__
#define __gm__
#endif
#ifndef __aicore__
#define __aicore__ [aicore]
#endif

#ifdef MEMORY_BASE
#undef MEMORY_BASE
#endif
#ifndef REGISTER_BASE
#define REGISTER_BASE
#endif

#include <pto/pto-inst.hpp>

#include "backend/urma/urma_completion_kernel.h"
#include "platform_comm/comm_context.h"
#include "tensor.h"

using namespace pto;

namespace {

constexpr int kElems = 128 * 128;

template <typename T>
static inline __aicore__ __gm__ T *tensor_data(__gm__ Tensor *tensor) {
return reinterpret_cast<__gm__ T *>(tensor->buffer.addr) + tensor->start_offset;
}

} // namespace

extern "C" __aicore__ __attribute__((always_inline)) void kernel_entry(__gm__ int64_t *args) {
__gm__ Tensor *input_tensor = reinterpret_cast<__gm__ Tensor *>(args[0]);
__gm__ Tensor *out_tensor = reinterpret_cast<__gm__ Tensor *>(args[1]);
__gm__ CommContext *comm_ctx = reinterpret_cast<__gm__ CommContext *>(args[2]);

// workSpace == 0 means the URMA overlay is not built in
// (SIMPLER_ENABLE_PTO_URMA_WORKSPACE=OFF, see docs/a5-sdma-overlay.md
// #1315): self-skip rather than dereferencing a null workspace.
if (comm_ctx == nullptr || comm_ctx->rankNum != 2 || comm_ctx->rankId >= comm_ctx->rankNum ||
comm_ctx->workSpace == 0 || comm_ctx->windowsIn[comm_ctx->rankId] == 0) {
pipe_barrier(PIPE_ALL);
return;
}

__gm__ float *local_input = tensor_data<float>(input_tensor);
__gm__ float *local_out = tensor_data<float>(out_tensor);
uint32_t peer_rank = 1u - comm_ctx->rankId;
uint64_t input_offset = reinterpret_cast<uint64_t>(local_input) - comm_ctx->windowsIn[comm_ctx->rankId];
__gm__ float *remote_input = pto2::urma_backend::peer_mr_ptr<float>(
reinterpret_cast<__gm__ uint8_t *>(comm_ctx->workSpace), peer_rank, input_offset
);

using FlatShape = Shape<1, 1, 1, 1, kElems>;
using FlatStride = pto::Stride<kElems, kElems, kElems, kElems, 1>;
using GlobalData = GlobalTensor<float, FlatShape, FlatStride>;

GlobalData remote_global(remote_input);
GlobalData local_global(local_out);

AsyncCtx async_ctx = get_async_ctx(args);
(void)send_request_entry(
async_ctx,
UrmaTget(local_global, remote_global, reinterpret_cast<__gm__ uint8_t *>(comm_ctx->workSpace), peer_rank)
);
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
/*
* 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.
* -----------------------------------------------------------------------------------------------------------
*/

#include <stdint.h>

#include "platform_comm/comm_context.h"
#include "pto_orchestration_api.h"

extern "C" {

__attribute__((visibility("default"))) PTO2OrchestrationConfig
urma_deferred_completion_orchestration_config(const L2TaskArgs &orch_args) {
(void)orch_args;
return PTO2OrchestrationConfig{.expected_arg_count = 4};
}

__attribute__((visibility("default"))) PTO2OrchestrationConfig aicpu_orchestration_config(const L2TaskArgs &orch_args) {
return urma_deferred_completion_orchestration_config(orch_args);
}

__attribute__((visibility("default"))) void urma_deferred_completion_orchestration(const L2TaskArgs &orch_args) {
if (orch_args.tensor_count() != 3 || orch_args.scalar_count() != 1) {
LOG_ERROR("urma_deferred_completion_demo: expected 3 tensors and 1 scalar");
return;
}

const Tensor &input = orch_args.tensor(0).ref();
const Tensor &out = orch_args.tensor(1).ref();
const Tensor &result = orch_args.tensor(2).ref();
auto *comm_ctx = reinterpret_cast<CommContext *>(static_cast<uintptr_t>(orch_args.scalar(0)));

L0TaskArgs producer_args;
producer_args.add_input(input);
producer_args.add_output(out);
producer_args.add_scalar(reinterpret_cast<uint64_t>(comm_ctx));
rt_submit_aiv_task(0, producer_args);

L0TaskArgs consumer_args;
consumer_args.add_input(out);
consumer_args.add_output(result);
rt_submit_aiv_task(1, consumer_args);
}

} // extern "C"
Loading