From 64259615303c17fdd72b8ae927d6ca80338a1d8d Mon Sep 17 00:00:00 2001 From: Vedant2005goyal Date: Mon, 20 Jul 2026 17:01:44 +0000 Subject: [PATCH] Enable restore_tracker usage in device kernels Replace the STL-based implementation of clad::restore_tracker with a device-compatible implementation. The previous implementation relied on std::vector and std::map, which are unavailable in device code, preventing restore_tracker from being used inside GPU kernels. --- include/clad/Differentiator/RestoreTracker.h | 42 ++++++++++++++++++++ test/CUDA/RestoreTrackerDevice.cu | 30 ++++++++++++++ 2 files changed, 72 insertions(+) create mode 100644 test/CUDA/RestoreTrackerDevice.cu diff --git a/include/clad/Differentiator/RestoreTracker.h b/include/clad/Differentiator/RestoreTracker.h index 22f0b7c49..608620a05 100644 --- a/include/clad/Differentiator/RestoreTracker.h +++ b/include/clad/Differentiator/RestoreTracker.h @@ -7,6 +7,12 @@ #include #include #include +#ifndef Max_Records +#define Max_Records 64 +#endif +#ifndef Max_Bytes +#define Max_Bytes 1024 +#endif namespace clad { @@ -21,6 +27,41 @@ namespace clad { /// clad::tape is not viable. class restore_tracker { // m_data consists of pairs of memory addresses and bitwise values +#ifdef __CUDACC__ + struct MetaData { + char* addr; + size_t size; + size_t off; + }; + MetaData m_meta[Max_Records]; + uint8_t m_buf[Max_Bytes]; + size_t m_cnt = 0, m_off = 0; + +public: + __host__ __device__ restore_tracker() = default; + + template __host__ __device__ void store(const T& val) { + for (size_t i = 0; i < m_cnt; ++i) + if (m_meta[i].addr == (char*)&val) + return; + + if (m_cnt >= Max_Records || m_off + sizeof(T) > Max_Bytes) { + // Clad restore_tracker GPU capacity exceeded. Try again with larger value + return; + } + + m_meta[m_cnt] = {(char*)&val, sizeof(T), m_off}; + std::memcpy(m_buf + m_off, &val, sizeof(T)); + m_off += sizeof(T); + m_cnt++; + } + + __host__ __device__ void restore() { + for (size_t i = 0; i < m_cnt; ++i) + std::memcpy(m_meta[i].addr, m_buf + m_meta[i].off, m_meta[i].size); + m_cnt = m_off = 0; + } +#else using RawMemory = std::vector; using Address = char*; std::map m_data; @@ -47,6 +88,7 @@ class restore_tracker { } m_data.clear(); } +#endif }; } // namespace clad diff --git a/test/CUDA/RestoreTrackerDevice.cu b/test/CUDA/RestoreTrackerDevice.cu new file mode 100644 index 000000000..ca9d9b081 --- /dev/null +++ b/test/CUDA/RestoreTrackerDevice.cu @@ -0,0 +1,30 @@ +// RUN: %cladclang_cuda -I%S/../../include --cuda-path=%cudapath \ +// RUN: --cuda-gpu-arch=%cudaarch %cudaldflags -o RestoreTrackerDevice.out %s +// +// RUN: %cudarun ./RestoreTrackerDevice.out | %filecheck_exec %s +// +// REQUIRES: cuda-runtime +#include "clad/Differentiator/Differentiator.h" +#include + +__global__ void test_kernel() { + clad::restore_tracker tracker; + + double val = 3.14; + tracker.store(val); + val = 0.0; + tracker.restore(); + + if (val == 3.14) { + printf("Working on device!\n"); + } +} + +int main() { + + test_kernel<<<1, 1>>>(); + cudaDeviceSynchronize(); + + // CHECK-EXEC: Working on device! + return 0; +}