Skip to content
Open
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
7 changes: 0 additions & 7 deletions csrc/ascend/batch_invariant_logp_ascend.asc
Original file line number Diff line number Diff line change
Expand Up @@ -307,10 +307,3 @@ std::vector<torch::Tensor> batch_invariant_logp_ascend_forward(torch::Tensor log
}
return {logp, lse};
}

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.def("batch_invariant_logp_ascend",
&batch_invariant_logp_ascend_forward,
"Batch-invariant selected-token log-probability (Ascend C forward)");
}
444 changes: 444 additions & 0 deletions csrc/ascend/distributed/deterministic_collective_ascend.asc

Large diffs are not rendered by default.

38 changes: 38 additions & 0 deletions csrc/ascend/ops_npu.asc
Original file line number Diff line number Diff line change
@@ -0,0 +1,38 @@
// SPDX-License-Identifier: Apache-2.0
// Copyright (c) 2026 RL-Kernel Contributors

// Single Python module entry point for the unified Ascend C extension.
// Keep PYBIND11_MODULE in one translation unit as new .asc kernels are added.

#include <vector>

#include <torch/extension.h>

std::vector<torch::Tensor> batch_invariant_logp_ascend_forward(
torch::Tensor logits, torch::Tensor target, int64_t ignore_index);

int64_t deterministic_collective_create(
torch::Tensor staging, int64_t world_size, int64_t rank);
void deterministic_collective_destroy(int64_t handle);
void deterministic_collective_stage(int64_t handle, torch::Tensor input);
void deterministic_collective_reduce(
int64_t handle, torch::Tensor gathered, torch::Tensor output, int64_t slice_offset);

PYBIND11_MODULE(TORCH_EXTENSION_NAME, m)
{
m.def("batch_invariant_logp_ascend",
&batch_invariant_logp_ascend_forward,
"Batch-invariant selected-token log-probability (Ascend C forward)");
m.def("deterministic_collective_create",
&deterministic_collective_create,
"Deterministic TP-invariant collective state (Ascend)");
m.def("deterministic_collective_destroy",
&deterministic_collective_destroy,
"Release a deterministic collective state (Ascend)");
m.def("deterministic_collective_stage",
&deterministic_collective_stage,
"Stage a tensor into the collective staging buffer (Ascend)");
m.def("deterministic_collective_reduce",
&deterministic_collective_reduce,
"Fixed-tree ordered reduction over gathered rank tensors (Ascend)");
}
13 changes: 13 additions & 0 deletions rl_engine/_C_npu.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -8,3 +8,16 @@ def batch_invariant_logp_ascend(
target: torch.Tensor,
ignore_index: int,
) -> list[torch.Tensor]: ...
def deterministic_collective_create(
staging: torch.Tensor,
world_size: int,
rank: int,
) -> int: ...
def deterministic_collective_destroy(handle: int) -> None: ...
def deterministic_collective_stage(handle: int, input: torch.Tensor) -> None: ...
def deterministic_collective_reduce(
handle: int,
gathered: torch.Tensor,
output: torch.Tensor,
slice_offset: int,
) -> None: ...
Loading
Loading