CoolFace
Modelpublic

replicate/megablocks

sourceHugging Faceapache-2.0updated 21d agoView on Hugging Face
0likes122downloads
torch_binding.cpp118 linesDownload Raw Back to torch-ext
1#include <torch/library.h>2 3#include "registration.h"4#include "torch_binding.h"5 6#include "new_cumsum.h"7#include "new_histogram.h"8#include "new_indices.h"9#include "new_replicate.h"10#include "new_sort.h"11 12#include "grouped_gemm/grouped_gemm.h"13 14// void exclusive_cumsum(torch::Tensor x, int dim, torch::Tensor out) {15torch::Tensor exclusive_cumsum_wrapper(torch::Tensor x, int64_t dim, torch::Tensor out) {16  megablocks::exclusive_cumsum(x, dim, out);17  return out;18}19 20// void inclusive_cumsum(torch::Tensor x, int dim, torch::Tensor out) {21torch::Tensor inclusive_cumsum_wrapper(torch::Tensor x, int64_t dim, torch::Tensor out) {22  megablocks::inclusive_cumsum(x, dim, out);23  return out;24}25 26// torch::Tensor histogram(torch::Tensor x, int num_bins);27torch::Tensor histogram_wrapper(torch::Tensor x, int64_t num_bins) {28  return megablocks::histogram(x, num_bins);29}30 31// void indices(torch::Tensor padded_bins,32//   int block_size,33//   int output_block_rows,34//   int output_block_columns,35//   torch::Tensor out);36torch::Tensor indices_wrapper(torch::Tensor padded_bins,37                               int64_t block_size,38                               int64_t output_block_rows,39                               int64_t output_block_columns,40                               torch::Tensor out) {41  megablocks::indices(padded_bins, block_size, output_block_rows, output_block_columns, out);42  return out;43}44 45 46 47// Forward pass: replicate values from x according to bin sizes48// void replicate_forward(torch::Tensor x,49//   torch::Tensor bins,50//   torch::Tensor out);51torch::Tensor replicate_forward_wrapper(torch::Tensor x, torch::Tensor bins, torch::Tensor out) {52  megablocks::replicate_forward(x, bins, out);53  return out;54}55 56// // Backward pass: reduce gradients back to bins using segmented reduction57// void replicate_backward(torch::Tensor grad,58//    torch::Tensor bins,59//    torch::Tensor out);60torch::Tensor replicate_backward_wrapper(torch::Tensor grad, torch::Tensor bins, torch::Tensor out) {61  megablocks::replicate_backward(grad, bins, out);62  return out;63}64 65// // Public interface function for radix sorting with indices66// void sort(torch::Tensor x,67//   int end_bit,68//   torch::Tensor x_out,69//   torch::Tensor iota_out);70torch::Tensor sort_wrapper(torch::Tensor x, int64_t end_bit, torch::Tensor x_out, torch::Tensor iota_out) {71  megablocks::sort(x, end_bit, x_out, iota_out);72  return x_out;73}74 75// GroupedGemm operation76torch::Tensor gmm(torch::Tensor a, torch::Tensor b, torch::Tensor c, torch::Tensor batch_sizes, bool trans_a, bool trans_b) {77  grouped_gemm::GroupedGemm(a, b, c, batch_sizes, trans_a, trans_b);78  return c;79}80 81// Reference implementation:82//83// m.def("exclusive_cumsum", &exclusive_cumsum, "batched exclusive cumsum.");84// m.def("histogram", &histogram, "even width histogram.");85// m.def("inclusive_cumsum", &inclusive_cumsum, "batched inclusive cumsum");86// m.def("indices", &indices, "indices construction for sparse matrix.");87// m.def("replicate_forward", &replicate_forward, "(fwd) replicate a vector dynamically.");88// m.def("replicate_backward", &replicate_backward, "(bwd) replicate a vector dynamically.");89// m.def("sort", &sort, "key/value sort.");90 91TORCH_LIBRARY_EXPAND(TORCH_EXTENSION_NAME, ops) {92  ops.def("exclusive_cumsum(Tensor x, int dim, Tensor(a!) out) -> Tensor(a!)");93  ops.impl("exclusive_cumsum", torch::kCUDA, &exclusive_cumsum_wrapper);94 95  ops.def("inclusive_cumsum(Tensor x, int dim, Tensor(a!) out) -> Tensor(a!)");96  ops.impl("inclusive_cumsum", torch::kCUDA, &inclusive_cumsum_wrapper);97 98  ops.def("histogram(Tensor x, int num_bins) -> Tensor");99  ops.impl("histogram", torch::kCUDA, &histogram_wrapper);100 101  ops.def("indices(Tensor padded_bins, int block_size, int output_block_rows, int output_block_columns, Tensor(a!) out) -> Tensor(a!)");102  ops.impl("indices", torch::kCUDA, &indices_wrapper);103 104  ops.def("replicate_forward(Tensor x, Tensor bins, Tensor(a!) out) -> Tensor(a!)");105  ops.impl("replicate_forward", torch::kCUDA, &replicate_forward_wrapper);106 107  ops.def("replicate_backward(Tensor grad, Tensor bins, Tensor(a!) out) -> Tensor(a!)");108  ops.impl("replicate_backward", torch::kCUDA, &replicate_backward_wrapper);109  110  ops.def("sort(Tensor x, int end_bit, Tensor x_out, Tensor iota_out) -> Tensor(x_out)");111  ops.impl("sort", torch::kCUDA, &sort_wrapper);112 113  // Register the gmm GroupedGemm operation114  ops.def("gmm(Tensor (a!) a, Tensor (b!) b, Tensor(c!) c, Tensor batch_sizes, bool trans_a, bool trans_b) -> Tensor(c!)");115  ops.impl("gmm", torch::kCUDA, &gmm);116}117 118REGISTER_EXTENSION(TORCH_EXTENSION_NAME)