forked from IQ.Lvbs/IQ.Pilot
IQ.Pilot Release Commit @ 0798119
This commit is contained in:
74
tinygrad_repo/extra/thunder/amd/include/pyutils/pyutils.cuh
Normal file
74
tinygrad_repo/extra/thunder/amd/include/pyutils/pyutils.cuh
Normal file
@@ -0,0 +1,74 @@
|
||||
#pragma once
|
||||
|
||||
#include "util.cuh"
|
||||
#include <pybind11/pybind11.h>
|
||||
|
||||
namespace kittens {
|
||||
namespace py {
|
||||
|
||||
template<typename T> struct from_object {
|
||||
static T make(pybind11::object obj) {
|
||||
return obj.cast<T>();
|
||||
}
|
||||
};
|
||||
template<ducks::gl::all GL> struct from_object<GL> {
|
||||
static GL make(pybind11::object obj) {
|
||||
// Check if argument is a torch.Tensor
|
||||
if (pybind11::hasattr(obj, "__class__") &&
|
||||
obj.attr("__class__").attr("__name__").cast<std::string>() == "Tensor") {
|
||||
|
||||
// Check if tensor is contiguous
|
||||
if (!obj.attr("is_contiguous")().cast<bool>()) {
|
||||
throw std::runtime_error("Tensor must be contiguous");
|
||||
}
|
||||
if (obj.attr("device").attr("type").cast<std::string>() == "cpu") {
|
||||
throw std::runtime_error("Tensor must be on CUDA device");
|
||||
}
|
||||
|
||||
// Get shape, pad with 1s if needed
|
||||
std::array<int, 4> shape = {1, 1, 1, 1};
|
||||
auto py_shape = obj.attr("shape").cast<pybind11::tuple>();
|
||||
size_t dims = py_shape.size();
|
||||
if (dims > 4) {
|
||||
throw std::runtime_error("Expected Tensor.ndim <= 4");
|
||||
}
|
||||
for (size_t i = 0; i < dims; ++i) {
|
||||
shape[4 - dims + i] = pybind11::cast<int>(py_shape[i]);
|
||||
}
|
||||
|
||||
// Get data pointer using data_ptr()
|
||||
uint64_t data_ptr = obj.attr("data_ptr")().cast<uint64_t>();
|
||||
|
||||
// Create GL object using make_gl
|
||||
return make_gl<GL>(data_ptr, shape[0], shape[1], shape[2], shape[3]);
|
||||
}
|
||||
throw std::runtime_error("Expected a torch.Tensor");
|
||||
}
|
||||
};
|
||||
|
||||
template<typename T> concept has_dynamic_shared_memory = requires(T t) { { t.dynamic_shared_memory() } -> std::convertible_to<int>; };
|
||||
|
||||
template<typename> struct trait;
|
||||
template<typename MT, typename T> struct trait<MT T::*> { using member_type = MT; using type = T; };
|
||||
template<typename> using object = pybind11::object;
|
||||
template<auto kernel, typename TGlobal> static void bind_kernel(auto m, auto name, auto TGlobal::*... member_ptrs) {
|
||||
m.def(name, [](object<decltype(member_ptrs)>... args) {
|
||||
TGlobal __g__ {from_object<typename trait<decltype(member_ptrs)>::member_type>::make(args)...};
|
||||
if constexpr (has_dynamic_shared_memory<TGlobal>) {
|
||||
int __dynamic_shared_memory__ = (int)__g__.dynamic_shared_memory();
|
||||
hipFuncSetAttribute((void *) kernel, hipFuncAttributeMaxDynamicSharedMemorySize, __dynamic_shared_memory__);
|
||||
kernel<<<__g__.grid(), __g__.block(), __dynamic_shared_memory__>>>(__g__);
|
||||
} else {
|
||||
kernel<<<__g__.grid(), __g__.block()>>>(__g__);
|
||||
}
|
||||
});
|
||||
}
|
||||
template<auto function, typename TGlobal> static void bind_function(auto m, auto name, auto TGlobal::*... member_ptrs) {
|
||||
m.def(name, [](object<decltype(member_ptrs)>... args) {
|
||||
TGlobal __g__ {from_object<typename trait<decltype(member_ptrs)>::member_type>::make(args)...};
|
||||
function(__g__);
|
||||
});
|
||||
}
|
||||
|
||||
} // namespace py
|
||||
} // namespace kittens
|
||||
@@ -0,0 +1,7 @@
|
||||
#pragma once
|
||||
|
||||
#include <torch/extension.h>
|
||||
|
||||
#define CHECK_CUDA(x) TORCH_CHECK(x.device().is_cuda(), #x " must be a CUDA tensor")
|
||||
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
|
||||
#define CHECK_INPUT(x) CHECK_CUDA(x); CHECK_CONTIGUOUS(x)
|
||||
18
tinygrad_repo/extra/thunder/amd/include/pyutils/util.cuh
Normal file
18
tinygrad_repo/extra/thunder/amd/include/pyutils/util.cuh
Normal file
@@ -0,0 +1,18 @@
|
||||
#pragma once
|
||||
|
||||
#include "../ops/ops.cuh"
|
||||
#include <iostream>
|
||||
|
||||
#define CHECK_CUDA_ERROR(val) check((val), #val, __FILE__, __LINE__)
|
||||
template <typename T>
|
||||
void check(T err, char const* const func, char const* const file,
|
||||
int const line)
|
||||
{
|
||||
if (err != hipSuccess)
|
||||
{
|
||||
std::cerr << "HIP Runtime Error at: " << file << ":" << line
|
||||
<< std::endl;
|
||||
std::cerr << hipGetErrorString(err) << " " << func << std::endl;
|
||||
//std::exit(EXIT_FAILURE);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user