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
78 changes: 44 additions & 34 deletions hat/backends/ffi/cuda/src/main/native/cpp/cuda_backend.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,19 @@ bool isCudaTileIR(const char *image, size_t len) {
std::memcmp(image, kMagic, sizeof(kMagic)) == 0;
}

// Tile code generation starts at sm_80 for both nvcc and NVRTC. sm_7x,
// including the nvcc default sm_75, does not generate Tile code.
// https://docs.nvidia.com/cuda/cuda-tile-cpp-api-reference/general_principles.html
void assertTileArchSupport(int major, int minor, const std::string &cudaPath) {
if (major >= 8) {
return;
}
std::cerr << "Tile compilation requires sm_80 or newer, got sm_"
<< major << minor << ". CUDA source saved to " << cudaPath
<< std::endl;
std::exit(1);
}

bool isLibjsigLoaded() {
static const bool loaded = [] {
void *handle = dlopen(nullptr, RTLD_LAZY);
Expand All @@ -71,11 +84,10 @@ void warnNvrtcRequiresLibjsig() {
}
warned = true;
std::cout << "[HAT] WARNING: libjsig not preloaded. NVRTC Tile "
<< "compilation loads LLVM into the JVM, which may hijack "
<< "SIGSEGV and break JDK implicit null-pointer checks. "
<< "Falling back to nvcc for all Tile kernels in this "
<< "process. This warning is issued once. To use NVRTC "
<< "Tile safely, launch with:"
<< "handlers see signals first and forward them with a "
<< "polluted state, breaking JVM signal handling. Falling "
<< "back to nvcc for Tile kernels. This warning is issued "
<< "once. To use NVRTC Tile, launch with:"
<< std::endl;
if (const char *javaHome = std::getenv("JAVA_HOME");
javaHome != nullptr && javaHome[0] != '\0') {
Expand Down Expand Up @@ -340,10 +352,8 @@ void CudaBackend::showDeviceInfo() {

std::cout << "> Using device 0: " << name << std::endl;

// get compute capabilities and the device name
int major = 0, minor = 0;
CUDA_CHECK(cuDeviceGetAttribute(&major, CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR, device), "cuDeviceGetAttribute");
CUDA_CHECK(cuDeviceGetAttribute(&minor, CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR, device), "cuDeviceGetAttribute");
queryComputeCapability(major, minor);
std::cout << "> GPU Device has major=" << major << " minor=" << minor << " compute capability" << std::endl;

int warpSize;
Expand All @@ -365,11 +375,15 @@ void CudaBackend::showDeviceInfo() {
((totalGlobalMem > static_cast<unsigned long long>(4) * 1024 * 1024 * 1024L) ? "YES" : "NO") << std::endl;
}

std::string CudaBackend::obtainSMVersion() {
int major = 0, minor = 0;
CUDA_CHECK(cuDeviceGetAttribute(&major,CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR, device), "cuDeviceGetAttribute");
CUDA_CHECK(cuDeviceGetAttribute(&minor,CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR, device), "cuDeviceGetAttribute");
return std::string("sm_").append(std::to_string(major)).append(std::to_string(minor));
void CudaBackend::queryComputeCapability(int &major, int &minor) {
CUDA_CHECK(cuDeviceGetAttribute(&major,
CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
device),
"cuDeviceGetAttribute");
CUDA_CHECK(cuDeviceGetAttribute(&minor,
CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
device),
"cuDeviceGetAttribute");
}

bool CudaBackend::useNvrtcCompiler(int typeModel) const {
Expand All @@ -386,8 +400,9 @@ bool CudaBackend::useNvrtcCompiler(int typeModel) const {
std::exit(1);
}
}
// Tile IR is JIT'd in-process with LLVM, which can replace the HotSpot
// SIGSEGV handler. SIMT NVRTC emits PTX and does not need libjsig.
// Tile NVRTC handlers see signals first and forward them with a
// polluted state, breaking JVM signal handling. Preload libjsig so
// JVM handlers run first. SIMT NVRTC does not need libjsig.
if (useNvrtc && typeModel > 0 && !isLibjsigLoaded()) {
warnNvrtcRequiresLibjsig();
return false;
Expand All @@ -408,13 +423,19 @@ CudaImage *CudaBackend::nvcc(const CudaSource *cudaSource) {
}
const std::string ptxPath = tmpFileName(time, localDirectory, suffix);
const std::string cudaPath = tmpFileName(time, localDirectory, ".cu");
cudaSource->write(cudaPath);

// Obtain the compute capability and SM version
std::string smVersion = obtainSMVersion();
int major = 0;
int minor = 0;
queryComputeCapability(major, minor);
if (cudaSource->typeModel() > 0) {
assertTileArchSupport(major, minor, cudaPath);
}
const std::string smVersion =
std::string("sm_") + std::to_string(major) + std::to_string(minor);

// compile the generated code
int pid;
cudaSource->write(cudaPath);
if ((pid = fork()) == 0) { //child
const auto cudaCompiler = "nvcc";
std::vector<std::string> command;
Expand All @@ -424,12 +445,11 @@ CudaImage *CudaBackend::nvcc(const CudaSource *cudaSource) {
command.push_back("--tilecubin");
command.push_back("--std=c++20");
command.push_back("--enable-tile");
command.push_back("-arch");
command.push_back(smVersion);
} else {
command.push_back("-ptx");
command.push_back("-Wno-deprecated-gpu-targets");
}
command.push_back("-arch");
command.push_back(smVersion);

command.push_back(cudaPath);
if (cudaSource->lineInfo()) {
Expand Down Expand Up @@ -483,19 +503,9 @@ CudaImage *CudaBackend::nvrtc(const CudaSource *cudaSource) {

int major = 0;
int minor = 0;
CUDA_CHECK(cuDeviceGetAttribute(&major,
CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MAJOR,
device),
"cuDeviceGetAttribute");
CUDA_CHECK(cuDeviceGetAttribute(&minor,
CU_DEVICE_ATTRIBUTE_COMPUTE_CAPABILITY_MINOR,
device),
"cuDeviceGetAttribute");
if (tile && major < 8) {
std::cerr << "NVRTC Tile compilation requires a GPU newer than sm_75, got sm_"
<< major << minor << ". CUDA source saved to " << cudaPath
<< std::endl;
std::exit(1);
queryComputeCapability(major, minor);
if (tile) {
assertTileArchSupport(major, minor, cudaPath);
}

nvrtcProgram program;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -177,10 +177,10 @@ class CudaQueue final : public Backend::Queue {
CUdevice device;
CUcontext context;
bool useNvrtcCompiler(int typeModel) const;
void queryComputeCapability(int &major, int &minor);
public:
void shortDeviceInfo() override;
void showDeviceInfo() override;
std::string obtainSMVersion();
CudaModule * compile(const CudaSource *cudaSource);
CudaModule * compile(const CudaSource &cudaSource);
CudaModule * compile(const CudaImage *image);
Expand Down