1 /*
2 * Copyright (c) 2024, 2026, Oracle and/or its affiliates. All rights reserved.
3 * DO NOT ALTER OR REMOVE COPYRIGHT NOTICES OR THIS FILE HEADER.
4 *
5 * This code is free software; you can redistribute it and/or modify it
6 * under the terms of the GNU General Public License version 2 only, as
7 * published by the Free Software Foundation. Oracle designates this
8 * particular file as subject to the "Classpath" exception as provided
9 * by Oracle in the LICENSE file that accompanied this code.
10 *
11 * This code is distributed in the hope that it will be useful, but WITHOUT
12 * ANY WARRANTY; without even the implied warranty of MERCHANTABILITY or
13 * FITNESS FOR A PARTICULAR PURPOSE. See the GNU General Public License
14 * version 2 for more details (a copy is included in the LICENSE file that
15 * accompanied this code).
16 *
17 * You should have received a copy of the GNU General Public License version
18 * 2 along with this work; if not, write to the Free Software Foundation,
19 * Inc., 51 Franklin St, Fifth Floor, Boston, MA 02110-1301 USA.
20 *
21 * Please contact Oracle, 500 Oracle Parkway, Redwood Shores, CA 94065 USA
22 * or visit www.oracle.com if you need additional information or have any
23 * questions.
24 */
25 #pragma once
26 #define CUDA_TYPES
27 #ifdef __APPLE__
28
29 #define LongUnsignedNewline "%llu\n"
30 #define Size_tNewline "%lu\n"
31 #define LongHexNewline "(0x%llx)\n"
32 #define alignedMalloc(size, alignment) memalign(alignment, size)
33 #define SNPRINTF snprintf
34 #else
35
36 #include <malloc.h>
37
38 #define LongHexNewline "(0x%lx)\n"
39 #define LongUnsignedNewline "%lu\n"
40 #define Size_tNewline "%lu\n"
41 #if defined (_WIN32)
42 #include "windows.h"
43 #define alignedMalloc(size, alignment) _aligned_malloc(size, alignment)
44 #define SNPRINTF _snprintf
45 #else
46 #define alignedMalloc(size, alignment) memalign(alignment, size)
47 #define SNPRINTF snprintf
48 #endif
49 #endif
50
51 #include <iostream>
52 #include <cuda.h>
53 #include <builtin_types.h>
54
55 #include "shared.h"
56
57 #include <fstream>
58 #include <thread>
59
60 struct WHERE{
61 const char* f;
62 int l;
63 cudaError_enum e;
64 const char* t;
65 void report() const {
66 if (e != CUDA_SUCCESS){
67 const char *buf;
68 cuGetErrorName(e, &buf);
69 std::cerr << t << " CUDA error = " << e << " " << buf <<std::endl<< " " << f << " line " << l << std::endl;
70 exit(-1);
71 }
72 }
73 };
74
75 #define CUDA_CHECK(err, functionName) { \
76 WHERE{.f =__FILE__, \
77 .l=__LINE__, \
78 .e = err, \
79 .t = functionName \
80 }.report(); \
81 }
82
83 // Loadable GPU image for cuModuleLoadData(Ex): PTX, cubin, or cuda_tile IR.
84 class CudaImage final : public Text {
85 public:
86 CudaImage();
87 explicit CudaImage(size_t len);
88 CudaImage(size_t len, char *text);
89 CudaImage(size_t len, char *text, bool isCopy);
90 explicit CudaImage(char *text);
91 ~CudaImage() override = default;
92 };
93
94 class CudaSource final :public Text {
95 public:
96 CudaSource(size_t len, char *text, bool isCopy, bool lineinfo, int typeModel);
97 bool lineInfo() const;
98 int typeModel() const;
99 explicit CudaSource(size_t len);
100 explicit CudaSource(char* text);
101 CudaSource();
102 ~CudaSource() override = default;
103 private:
104 bool _lineInfo = false;
105 int _typeModel = 0;
106 };
107
108 class CudaBackend final : public Backend {
109 public:
110 class CudaQueue final : public Backend::Queue {
111 public:
112 std::thread::id streamCreationThread;
113 CUstream cuStream;
114 explicit CudaQueue(Backend *backend);
115 void init();
116 void wait() override;
117
118 void release() override;
119
120 void computeStart() override;
121
122 void computeEnd() override;
123
124 void copyToDevice(Buffer *buffer) override;
125
126 void copyFromDevice(Buffer *buffer) override;
127
128 int estimateThreadsPerBlock(int dimensions);
129
130 int estimateThreadsPerBlock(int dimensions, int globalSizePerDimension, int localSize);
131
132 void dispatch(DispatchContext *dispatchContext, CompilationUnit::Kernel *kernel) override;
133
134 ~CudaQueue() override;
135 };
136
137 class CudaBuffer final : public Buffer {
138 public:
139 CUdeviceptr devicePtr;
140 CudaBuffer(Backend *backend, BufferState *bufferState);
141 ~CudaBuffer() override;
142 };
143
144 class CudaModule final : public CompilationUnit {
145 CUmodule module;
146 CudaSource cudaSource;
147 CudaImage image;
148 Log log;
149
150 public:
151 class CudaKernel final : public Kernel {
152
153 public:
154 bool setArg(KernelArg *arg) override;
155 bool setArg(KernelArg *arg, Buffer *buffer) override;
156 CudaKernel(Backend::CompilationUnit *program, char* name, CUfunction function);
157 ~CudaKernel() override;
158 static CudaKernel * of(long kernelHandle);
159 static CudaKernel * of(Backend::CompilationUnit::Kernel *kernel);
160
161 CUfunction function;
162 void *argslist[100]{};
163 };
164 CudaModule(Backend *backend, const CudaImage *image, char *log,
165 bool ok, CUmodule module);
166 ~CudaModule() override;
167 static CudaModule * of(long moduleHandle);
168 //static CudaModule * of(CompilationUnit *compilationUnit);
169 Kernel *getKernel(int nameLen, char *name) override;
170 CudaKernel *getCudaKernel(char *name);
171 CudaKernel *getCudaKernel(int nameLen, char *name);
172 bool programOK();
173 };
174
175 private:
176 CUresult initStatus;
177 CUdevice device;
178 CUcontext context;
179 bool useNvrtcCompiler(int typeModel) const;
180 public:
181 void shortDeviceInfo() override;
182 void showDeviceInfo() override;
183 std::string obtainSMVersion();
184 CudaModule * compile(const CudaSource *cudaSource);
185 CudaModule * compile(const CudaSource &cudaSource);
186 CudaModule * compile(const CudaImage *image);
187 CudaModule * compile(const CudaImage &image);
188 CudaImage *nvcc(const CudaSource *cudaSource);
189 CudaImage *nvrtc(const CudaSource *cudaSource);
190 CompilationUnit * compile(int len, char *source, int typeModel) override;
191 void computeStart() override;
192 void computeEnd() override;
193 CudaBuffer * getOrCreateBuffer(BufferState *bufferState) override;
194 bool getBufferFromDeviceIfDirty(void *memorySegment, long memorySegmentLength) override;
195
196 explicit CudaBackend(int mode);
197
198 ~CudaBackend() override;
199 static CudaBackend * of(long backendHandle);
200 static CudaBackend * of(Backend *backend);
201 };