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 };