From 8ecbb7bf991984fedb0cba920ffcd4055b6ba0ef Mon Sep 17 00:00:00 2001 From: Christian Stewart Date: Tue, 25 Aug 2026 19:24:38 -0700 Subject: [PATCH 1/3] Add HIP backend for the fused PPO trainer on AMD ROCm Translate the CUDA training sources to HIP with HIPIFY and build them with hipcc, so the fused trainer runs on AMD GPUs. The CUDA/cuDNN build path is unchanged; everything new is additive. - src-hip/ holds the hipified bindings, kernels, models, muon, ocean, and pufferlib sources, plus a cuda->HIP forwarding shim used by the C env libraries and a stub NVML header (utilization reads as zeros). - The conv encoder path has dual implementations: NVIDIA builds keep src/cudnn_conv2d.cu (cuDNN) untouched via ocean.cu, while AMD builds use src-hip/cudnn_conv2d.hip.cpp (MIOpen) with an identical interface via ocean.hip.cpp. Forward is miopenConvolutionForward with bias + in-place ReLU; backward uses the MIOpen weights/data kernels. The nethack encoder is not ported. - src/vecenv.h maps its small CUDA surface onto HIP equivalents under __HIP_PLATFORM_AMD__; other platforms keep the existing declarations. - build-hip.sh mirrors build.sh: it compiles the static env library, then builds pufferlib/_C with hipcc --offload-arch= linking amdhip64, hipblas, hiprand, rocrand, rccl, and MIOpen. Verified on an RX 7700 XT (gfx1101) with ROCm 7.2.2: cartpole trains through the fused rollout path at ~6M SPS. --- build-hip.sh | 52 + src-hip/bindings.hip.cpp | 635 +++++++++ src-hip/cuda_shim.cpp | 18 + src-hip/cudnn_conv2d.hip.cpp | 231 ++++ src-hip/kernels.hip.cpp | 470 +++++++ src-hip/models.hip.cpp | 855 +++++++++++++ src-hip/muon.hip.cpp | 228 ++++ src-hip/ocean.hip.cpp | 599 +++++++++ src-hip/ocean_stub.cpp | 1 + src-hip/pufferlib.hip.cpp | 2337 ++++++++++++++++++++++++++++++++++ src-hip/stub/nvml.h | 20 + src/vecenv.h | 23 + 12 files changed, 5469 insertions(+) create mode 100755 build-hip.sh create mode 100644 src-hip/bindings.hip.cpp create mode 100644 src-hip/cuda_shim.cpp create mode 100644 src-hip/cudnn_conv2d.hip.cpp create mode 100644 src-hip/kernels.hip.cpp create mode 100644 src-hip/models.hip.cpp create mode 100644 src-hip/muon.hip.cpp create mode 100644 src-hip/ocean.hip.cpp create mode 100644 src-hip/ocean_stub.cpp create mode 100644 src-hip/pufferlib.hip.cpp create mode 100644 src-hip/stub/nvml.h diff --git a/build-hip.sh b/build-hip.sh new file mode 100755 index 0000000000..827bbc1c3c --- /dev/null +++ b/build-hip.sh @@ -0,0 +1,52 @@ +#!/bin/bash +# build-hip.sh - AMD ROCm fused trainer (_C.so) via HIPIFY-translated sources +set -e +ENV=${1:?usage: build-hip.sh ENV [cartpole|breakout|...]} +cd "$(dirname "$0")" + +SRC_DIR="ocean/$ENV" +RAYLIB_NAME='raylib-5.5_linux_amd64' +if [ ! -d "$RAYLIB_NAME" ]; then + curl -sL "https://github.com/raysan5/raylib/releases/download/5.5/$RAYLIB_NAME.tar.gz" -o r.tgz + tar xf r.tgz && rm r.tgz +fi + +HIPCC=/opt/rocm/bin/hipcc +ARCH=${HIP_ARCH:-gfx1101} +ROCM=/opt/rocm + +PYTHON_INCLUDE=$(python -c "import sysconfig; print(sysconfig.get_path('include'))") +PYBIND_INCLUDE=$(python -c "import pybind11; print(pybind11.get_include())") +NUMPY_INCLUDE=$(python -c "import numpy; print(numpy.get_include())") +EXT_SUFFIX=$(python -c "import sysconfig; print(sysconfig.get_config_var('EXT_SUFFIX'))") +OUTPUT="pufferlib/_C${EXT_SUFFIX}" + +BINDING_SRC="$SRC_DIR/binding.c" +mkdir -p build +STATIC_LIB="build/libstatic_${ENV}.a" +clang -c -O2 -std=gnu11 -D_GNU_SOURCE -I. -Isrc -I"$SRC_DIR" -Ivendor \ + -I./$RAYLIB_NAME/include -DPLATFORM_DESKTOP \ + -fno-semantic-interposition -fvisibility=hidden -fPIC -fopenmp \ + "$BINDING_SRC" -o build/libstatic_${ENV}.o +ar rcs "$STATIC_LIB" build/libstatic_${ENV}.o +OBS_TENSOR_T=$(awk '/^#define OBS_TENSOR_T/{print $3}' "$BINDING_SRC") + +echo "Compiling HIP training backend ($ARCH)..." +"$HIPCC" -c --offload-arch=$ARCH -fPIC -std=c++17 -O2 \ + -I. -Isrc -Isrc-hip -Isrc-hip/stub \ + -I"$PYTHON_INCLUDE" -I"$PYBIND_INCLUDE" -I"$NUMPY_INCLUDE" \ + -I$ROCM/include -I$ROCM/include/hipblas -I$ROCM/include/hiprand -I./$RAYLIB_NAME/include \ + -D_GLIBCXX_USE_CXX11_ABI=1 -DNPY_NO_DEPRECATED_API=NPY_1_7_API_VERSION \ + -DPLATFORM_DESKTOP -fopenmp \ + -DOBS_TENSOR_T=$OBS_TENSOR_T -DENV_NAME=$ENV -DPRECISION_FLOAT \ + src-hip/bindings.hip.cpp -o build/bindings_hip.o + +"$HIPCC" -c -fPIC -std=c++17 -O2 -x hip --offload-arch=$ARCH \ + src-hip/cuda_shim.cpp -o build/cuda_shim.o +g++ -shared -fPIC -fopenmp \ + build/bindings_hip.o build/cuda_shim.o "$STATIC_LIB" "$RAYLIB_NAME/lib/libraylib.a" \ + -L$ROCM/lib -lamdhip64 -lhipblas -lhiprand -lrocrand -lrccl \ + -Wl,-rpath,$ROCM/lib -Bsymbolic-functions \ + -lm -lpthread -lomp5 \ + -o "$OUTPUT" +echo "Built: $OUTPUT" diff --git a/src-hip/bindings.hip.cpp b/src-hip/bindings.hip.cpp new file mode 100644 index 0000000000..ca280f8e83 --- /dev/null +++ b/src-hip/bindings.hip.cpp @@ -0,0 +1,635 @@ +#include "hip/hip_runtime.h" +// bindings.cpp - Python bindings for pufferlib (torch-free) + +#include +#include +#include +#include "pufferlib.hip.cpp" + +#define _PUFFER_STRINGIFY(x) #x +#define PUFFER_STRINGIFY(x) _PUFFER_STRINGIFY(x) + +namespace py = pybind11; + +// Wrapper functions for Python bindings +pybind11::dict puf_log(pybind11::object pufferl_obj) { + auto& pufferl = pufferl_obj.cast(); + pybind11::dict result; + + // Summary + int gpus = pufferl.hypers.world_size; + long global_step = pufferl.global_step; + long epoch = pufferl.epoch; + double now = wall_clock(); + double dt = now - pufferl.last_log_time; + long sps = dt > 0 ? (long)((global_step - pufferl.last_log_step) / dt) : 0; + pufferl.last_log_time = now; + pufferl.last_log_step = global_step; + + result["SPS"] = sps * gpus; + result["agent_steps"] = global_step * gpus; + result["uptime"] = now - pufferl.start_time; + result["epoch"] = epoch; + + // Environment stats + pybind11::dict env_dict; + Dict* env_out = log_environments_impl(pufferl); + for (int i = 0; i < env_out->size; i++) { + env_dict[env_out->items[i].key] = env_out->items[i].value; + } + result["env"] = env_dict; + + // Losses + pybind11::dict losses_dict; + float losses_host[NUM_LOSSES]; + hipMemcpy(losses_host, pufferl.losses_puf.data, sizeof(losses_host), hipMemcpyDeviceToHost); + float n = losses_host[LOSS_N]; + if (n > 0) { + float inv_n = 1.0f / n; + losses_dict["policy"] = losses_host[LOSS_PG] * inv_n; + losses_dict["value"] = losses_host[LOSS_VF] * inv_n; + losses_dict["entropy"] = losses_host[LOSS_ENT] * inv_n; + losses_dict["total"] = losses_host[LOSS_TOTAL] * inv_n; + losses_dict["old_kl"] = losses_host[LOSS_OLD_APPROX_KL] * inv_n; + losses_dict["kl"] = losses_host[LOSS_APPROX_KL] * inv_n; + losses_dict["clipfrac"] = losses_host[LOSS_CLIPFRAC] * inv_n; + } + hipMemset(pufferl.losses_puf.data, 0, numel(pufferl.losses_puf.shape) * sizeof(float)); + result["loss"] = losses_dict; + + // Profile + pybind11::dict perf_dict; + float train_total = 0; + for (int i = 0; i < NUM_PROF; i++) { + float sec = pufferl.profile.accum[i] / 1000.0f; + perf_dict[PROF_NAMES[i]] = sec; + if (i >= PROF_TRAIN_MISC) train_total += sec; + } + perf_dict["train"] = train_total; + memset(pufferl.profile.accum, 0, sizeof(pufferl.profile.accum)); + result["perf"] = perf_dict; + + // Utilization + pybind11::dict util_dict; + nvmlUtilization_t util; + nvmlDeviceGetUtilizationRates(pufferl.nvml_device, &util); + util_dict["gpu_percent"] = (float)util.gpu; + + nvmlMemory_t mem; + nvmlDeviceGetMemoryInfo(pufferl.nvml_device, &mem); + util_dict["gpu_mem"] = 100.0f * (float)mem.used / (float)mem.total; + + size_t cuda_free, cuda_total; + hipMemGetInfo(&cuda_free, &cuda_total); + util_dict["vram_used_gb"] = (float)(cuda_total - cuda_free) / (1024.0f * 1024.0f * 1024.0f); + util_dict["vram_total_gb"] = (float)cuda_total / (1024.0f * 1024.0f * 1024.0f); + + long rss_kb = 0; + FILE* f = fopen("/proc/self/status", "r"); + if (f) { + char line[256]; + while (fgets(line, sizeof(line), f)) { + if (sscanf(line, "VmRSS: %ld", &rss_kb) == 1) break; + } + fclose(f); + } + util_dict["cpu_mem_gb"] = (float)rss_kb / (1024.0f * 1024.0f); + result["util"] = util_dict; + + return result; +} + +pybind11::dict puf_eval_log(pybind11::object pufferl_obj) { + auto& pufferl = pufferl_obj.cast(); + pybind11::dict result; + + double now = wall_clock(); + pufferl.last_log_time = now; + pufferl.last_log_step = pufferl.global_step; + + pybind11::dict env_dict; + // Capacity 64 to fit chess's per-bank hist_score_bank/hist_n_bank entries + // (16 keys across 8 banks) on top of base env-log fields. + Dict* env_out = create_dict(64); + static_vec_eval_log(pufferl.vec, env_out); + for (int i = 0; i < env_out->size; i++) { + env_dict[env_out->items[i].key] = env_out->items[i].value; + } + result["env"] = env_dict; + + return result; +} + +void python_vec_recv(pybind11::object pufferl_obj, int buf) { + // Not used in static/OMP path +} + +void python_vec_send(pybind11::object pufferl_obj, int buf) { + // Not used in static/OMP path +} + +void render(pybind11::object pufferl_obj, int env_id) { + PuffeRL& pufferl = pufferl_obj.cast(); + static_vec_render(pufferl.vec, env_id); +} + +void rollouts(pybind11::object pufferl_obj) { + PuffeRL& pufferl = pufferl_obj.cast(); + pybind11::gil_scoped_release no_gil; + double t0 = wall_clock(); + + // Zero state buffers (primary + every frozen bank, so all banks see fresh + // state symmetrically — otherwise frozen banks accumulate indefinitely while + // primary resets, giving primary an unfair in-distribution advantage). + if (pufferl.hypers.reset_state) { + for (int i = 0; i < pufferl.hypers.num_buffers; i++) { + puf_zero(&pufferl.buffer_states[i], pufferl.default_stream); + } + for (int b = 0; b < pufferl.num_frozen_banks; b++) { + for (int i = 0; i < pufferl.hypers.num_buffers; i++) { + puf_zero(&pufferl.frozen_banks[b].buffer_states[i], pufferl.default_stream); + } + } + } + + static_vec_omp_step(pufferl.vec); + float sec = (float)(wall_clock() - t0); + pufferl.profile.accum[PROF_ROLLOUT] += sec * 1000.0f; // store as ms + + float eval_prof[NUM_EVAL_PROF]; + static_vec_read_profile(pufferl.vec, eval_prof); + pufferl.profile.accum[PROF_EVAL_GPU] += eval_prof[EVAL_GPU]; + pufferl.profile.accum[PROF_EVAL_ENV] += eval_prof[EVAL_ENV_STEP]; + pufferl.global_step += pufferl.hypers.horizon * pufferl.hypers.total_agents; +} + +pybind11::dict train(pybind11::object pufferl_obj) { + PuffeRL& pufferl = pufferl_obj.cast(); + { + pybind11::gil_scoped_release no_gil; + train_impl(pufferl); + } + pybind11::dict losses; + return losses; +} + +void puf_close(pybind11::object pufferl_obj) { + PuffeRL& pufferl = pufferl_obj.cast(); + close_impl(pufferl); +} + +void save_weights(pybind11::object pufferl_obj, const std::string& path) { + PuffeRL& pufferl = pufferl_obj.cast(); + int64_t nbytes = numel(pufferl.master_weights.shape) * sizeof(float); + std::vector buf(nbytes); + hipMemcpy(buf.data(), pufferl.master_weights.data, nbytes, hipMemcpyDeviceToHost); + FILE* f = fopen(path.c_str(), "wb"); + if (!f) throw std::runtime_error("Failed to open " + path + " for writing"); + fwrite(buf.data(), 1, nbytes, f); + fclose(f); +} + +void load_weights(pybind11::object pufferl_obj, const std::string& path) { + PuffeRL& pufferl = pufferl_obj.cast(); + int64_t nbytes = numel(pufferl.master_weights.shape) * sizeof(float); + FILE* f = fopen(path.c_str(), "rb"); + if (!f) throw std::runtime_error("Failed to open " + path + " for reading"); + // Verify file size matches + fseek(f, 0, SEEK_END); + long file_size = ftell(f); + fseek(f, 0, SEEK_SET); + if (file_size != nbytes) { + fclose(f); + throw std::runtime_error("Weight file size mismatch: expected " + + std::to_string(nbytes) + " bytes, got " + std::to_string(file_size)); + } + std::vector buf(nbytes); + size_t nread = fread(buf.data(), 1, nbytes, f); + if ((int64_t)nread != nbytes) { + fclose(f); + throw std::runtime_error("Failed to read weight file"); + } + fclose(f); + hipMemcpy(pufferl.master_weights.data, buf.data(), nbytes, hipMemcpyHostToDevice); + if (USE_BF16) { + int n = numel(pufferl.param_puf.shape); + cast<<>>( + pufferl.param_puf.data, pufferl.master_weights.data, n); + } +} + +int py_add_frozen_bank(py::object pufferl_obj, int slice_size, + int hidden_size, int num_layers) { + PuffeRL& pufferl = pufferl_obj.cast(); + return pufferl_add_frozen_bank(&pufferl, slice_size, hidden_size, num_layers); +} + +void py_load_frozen_bank(py::object pufferl_obj, int bank_idx, const std::string& path) { + PuffeRL& pufferl = pufferl_obj.cast(); + pufferl_load_frozen_bank(&pufferl, bank_idx, path.c_str()); +} + +void py_set_agent_perm(py::object pufferl_obj, py::array_t perm) { + PuffeRL& pufferl = pufferl_obj.cast(); + auto buf = perm.request(); + if (buf.ndim != 1) throw std::runtime_error("agent_perm must be 1-D"); + if ((int)buf.shape[0] != pufferl.vec->total_agents) { + throw std::runtime_error("agent_perm length must equal total_agents"); + } + pufferl_set_agent_perm(&pufferl, (const int*)buf.ptr); +} + +void py_set_env_tags(py::object pufferl_obj, py::array_t tags) { + PuffeRL& pufferl = pufferl_obj.cast(); + auto buf = tags.request(); + if (buf.ndim != 1) throw std::runtime_error("env_tags must be 1-D"); + int num_envs = pufferl_num_envs(&pufferl); + if ((int)buf.shape[0] != num_envs) { + throw std::runtime_error("env_tags length must equal num_envs"); + } + pufferl_set_env_tags(&pufferl, (const int*)buf.ptr); +} + +int py_count_aligned(py::object pufferl_obj, int tag_value, int reset_flags) { + PuffeRL& pufferl = pufferl_obj.cast(); + return pufferl_count_aligned(&pufferl, tag_value, reset_flags); +} + +int py_num_envs(py::object pufferl_obj) { + PuffeRL& pufferl = pufferl_obj.cast(); + return pufferl_num_envs(&pufferl); +} + +void py_puff_advantage( + long long values_ptr, long long rewards_ptr, + long long dones_ptr, long long importance_ptr, + long long advantages_ptr, + int num_steps, int horizon, + float gamma, float lambda, float rho_clip, float c_clip) { + constexpr int N = 16 / sizeof(precision_t); + int blocks = grid_size(num_steps); + auto kernel = (horizon % N == 0) ? puff_advantage : puff_advantage_scalar; + kernel<<>>( + (const precision_t*)values_ptr, (const precision_t*)rewards_ptr, + (const precision_t*)dones_ptr, (const precision_t*)importance_ptr, + (precision_t*)advantages_ptr, + gamma, lambda, rho_clip, c_clip, num_steps, horizon); +} + +double get_config(py::dict& kwargs, const char* key) { + if (!kwargs.contains(key)) { + throw std::runtime_error(std::string("Missing config key: ") + key); + } + try { + return kwargs[key].cast(); + } catch (const py::cast_error& e) { + throw std::runtime_error(std::string("Failed to cast config key '") + key + "': " + e.what()); + } +} + +Dict* py_dict_to_c_dict(py::dict py_dict) { + Dict* c_dict = create_dict(py_dict.size()); + for (auto item : py_dict) { + const char* key = PyUnicode_AsUTF8(item.first.ptr()); + try { + dict_set(c_dict, key, item.second.cast()); + } catch (const py::cast_error&) { + // Skip non-numeric values + } + } + return c_dict; +} + +// ============================================================================ +// Python-facing VecEnv: wraps StaticVec for use from python_pufferl.py. +// After vec_step(), GPU buffers are current — Python wraps them zero-copy +// with torch.from_blob(ptr, shape, dtype, device='cuda'). +// ============================================================================ + +struct VecEnv { + StaticVec* vec; + int total_agents; + int obs_size; + int num_atns; + std::vector act_sizes; + std::string obs_dtype; + size_t obs_elem_size; + int gpu; +}; + +std::unique_ptr create_vec(py::dict args, int gpu) { + py::dict vec_kwargs = args["vec"].cast(); + py::dict env_kwargs = args["env"].cast(); + + int total_agents = (int)get_config(vec_kwargs, "total_agents"); + int num_buffers = (int)get_config(vec_kwargs, "num_buffers"); + + Dict* vec_dict = py_dict_to_c_dict(vec_kwargs); + Dict* env_dict = py_dict_to_c_dict(env_kwargs); + + auto ve = std::make_unique(); + ve->gpu = gpu; + { + py::gil_scoped_release no_gil; + ve->vec = create_static_vec(total_agents, num_buffers, gpu, vec_dict, env_dict); + } + ve->total_agents = total_agents; + ve->obs_size = get_obs_size(); + ve->num_atns = get_num_atns(); + { + int* raw = get_act_sizes(); + int n = get_num_act_sizes(); + ve->act_sizes = std::vector(raw, raw + n); + } + ve->obs_dtype = std::string(get_obs_dtype()); + ve->obs_elem_size = get_obs_elem_size(); + return ve; +} + +void vec_reset(VecEnv& ve) { + py::gil_scoped_release no_gil; + static_vec_reset(ve.vec); +} + +void gpu_vec_step_py(VecEnv& ve, long long actions_ptr) { + hipMemcpy(ve.vec->gpu_actions, (void*)actions_ptr, + (size_t)ve.total_agents * ve.num_atns * sizeof(float), + hipMemcpyDeviceToDevice); + { + py::gil_scoped_release no_gil; + gpu_vec_step(ve.vec); + } +} + +void cpu_vec_step_py(VecEnv& ve, long long actions_ptr) { + memcpy(ve.vec->actions, (void*)actions_ptr, + (size_t)ve.total_agents * ve.num_atns * sizeof(float)); + { + py::gil_scoped_release no_gil; + cpu_vec_step(ve.vec); + } +} + +py::dict vec_log(VecEnv& ve) { + Dict* out = create_dict(32); + static_vec_log(ve.vec, out); + py::dict result; + for (int i = 0; i < out->size; i++) { + result[out->items[i].key] = out->items[i].value; + } + free(out->items); + free(out); + return result; +} + +void vec_close(VecEnv& ve) { + static_vec_close(ve.vec); + ve.vec = nullptr; +} + +std::unique_ptr create_pufferl(py::dict args) { + py::dict train_kwargs = args["train"].cast(); + py::dict vec_kwargs = args["vec"].cast(); + py::dict env_kwargs = args["env"].cast(); + py::dict policy_kwargs = args["policy"].cast(); + + HypersT hypers; + // Layout (total_agents and num_buffers come from vec config) + hypers.total_agents = get_config(vec_kwargs, "total_agents"); + hypers.num_buffers = get_config(vec_kwargs, "num_buffers"); + hypers.num_threads = get_config(vec_kwargs, "num_threads"); + hypers.horizon = get_config(train_kwargs, "horizon"); + // Model architecture (num_atns computed from env in C++) + hypers.hidden_size = get_config(policy_kwargs, "hidden_size"); + hypers.num_layers = get_config(policy_kwargs, "num_layers"); + // Learning rate + hypers.lr = get_config(train_kwargs, "learning_rate"); + hypers.min_lr_ratio = get_config(train_kwargs, "min_lr_ratio"); + hypers.anneal_lr = get_config(train_kwargs, "anneal_lr"); + // Optimizer + hypers.beta1 = get_config(train_kwargs, "beta1"); + hypers.beta2 = get_config(train_kwargs, "beta2"); + hypers.eps = get_config(train_kwargs, "eps"); + // Training + hypers.minibatch_size = get_config(train_kwargs, "minibatch_size"); + hypers.replay_ratio = get_config(train_kwargs, "replay_ratio"); + hypers.total_timesteps = get_config(train_kwargs, "total_timesteps"); + hypers.max_grad_norm = get_config(train_kwargs, "max_grad_norm"); + // PPO + hypers.clip_coef = get_config(train_kwargs, "clip_coef"); + hypers.vf_clip_coef = get_config(train_kwargs, "vf_clip_coef"); + hypers.vf_coef = get_config(train_kwargs, "vf_coef"); + hypers.ent_coef = get_config(train_kwargs, "ent_coef"); + hypers.min_ent_coef_ratio = get_config(train_kwargs, "min_ent_coef_ratio"); + hypers.anneal_ent_coef = get_config(train_kwargs, "anneal_ent_coef"); + // GAE + hypers.gamma = get_config(train_kwargs, "gamma"); + hypers.gae_lambda = get_config(train_kwargs, "gae_lambda"); + // VTrace + hypers.vtrace_rho_clip = get_config(train_kwargs, "vtrace_rho_clip"); + hypers.vtrace_c_clip = get_config(train_kwargs, "vtrace_c_clip"); + // Priority + hypers.prio_alpha = get_config(train_kwargs, "prio_alpha"); + hypers.prio_beta0 = get_config(train_kwargs, "prio_beta0"); + hypers.reset_state = get_config(args, "reset_state"); + // Base-level config ([base] section becomes top-level in args) + hypers.cudagraphs = get_config(args, "cudagraphs"); + hypers.profile = get_config(args, "profile"); + // Multi-GPU / device selection + hypers.rank = get_config(args, "rank"); + hypers.world_size = get_config(args, "world_size"); + hypers.gpu_id = get_config(args, "gpu_id"); + hypers.nccl_id = args["nccl_id"].cast(); + // Seed + hypers.seed = get_config(args, "seed"); + + int device_count = 0; + int err = hipGetDeviceCount(&device_count); + if (err != hipSuccess) { + throw std::runtime_error("CUDA is not available"); + } + + std::string env_name = args["env_name"].cast(); + Dict* vec_dict = py_dict_to_c_dict(vec_kwargs.cast()); + Dict* env_dict = py_dict_to_c_dict(env_kwargs.cast()); + + std::unique_ptr pufferl; + { + pybind11::gil_scoped_release no_gil; + pufferl = create_pufferl_impl(hypers, env_name, vec_dict, env_dict); + } + + if (!pufferl) { + throw std::runtime_error("CUDA OOM: failed to allocate training buffers"); + } + + return pufferl; +} + +PYBIND11_MODULE(_C, m) { + // Multi-GPU: generate NCCL unique ID (call on rank 0, pass bytes to all ranks) + m.def("get_nccl_id", []() { + ncclUniqueId id; + ncclGetUniqueId(&id); + return py::bytes(reinterpret_cast(&id), sizeof(id)); + }); + // Standalone utilization monitor (no PuffeRL instance needed) + m.def("get_utilization", [](int gpu_id) { + static bool nvml_inited = false; + if (!nvml_inited) { nvmlInit(); nvml_inited = true; } + + py::dict util_dict; + nvmlDevice_t device; + nvmlDeviceGetHandleByIndex(gpu_id, &device); + + nvmlUtilization_t util; + nvmlDeviceGetUtilizationRates(device, &util); + util_dict["gpu_percent"] = (float)util.gpu; + + nvmlMemory_t mem; + nvmlDeviceGetMemoryInfo(device, &mem); + util_dict["gpu_mem"] = 100.0f * (float)mem.used / (float)mem.total; + + size_t cuda_free, cuda_total; + hipMemGetInfo(&cuda_free, &cuda_total); + util_dict["vram_used_gb"] = (float)(cuda_total - cuda_free) / (1024.0f * 1024.0f * 1024.0f); + util_dict["vram_total_gb"] = (float)cuda_total / (1024.0f * 1024.0f * 1024.0f); + + long rss_kb = 0; + FILE* f = fopen("/proc/self/status", "r"); + if (f) { + char line[256]; + while (fgets(line, sizeof(line), f)) { + if (sscanf(line, "VmRSS: %ld", &rss_kb) == 1) break; + } + fclose(f); + } + util_dict["cpu_mem_gb"] = (float)rss_kb / (1024.0f * 1024.0f); + + return util_dict; + }); + + m.attr("precision_bytes") = (int)sizeof(precision_t); + m.attr("env_name") = PUFFER_STRINGIFY(ENV_NAME); + m.attr("gpu") = 1; + + // Core functions + m.def("log", &puf_log); + m.def("eval_log", &puf_eval_log); + m.def("render", &render); + m.def("rollouts", &rollouts); + m.def("train", &train); + m.def("close", &puf_close); + m.def("save_weights", &save_weights); + m.def("load_weights", &load_weights); + m.def("add_frozen_bank", &py_add_frozen_bank); + m.def("load_frozen_bank", &py_load_frozen_bank); + m.def("set_agent_perm", &py_set_agent_perm); + m.def("set_env_tags", &py_set_env_tags); + m.def("count_aligned", &py_count_aligned); + m.def("num_envs", &py_num_envs); + m.def("python_vec_recv", &python_vec_recv); + m.def("python_vec_send", &python_vec_send); + py::class_(m, "Policy"); + py::class_(m, "Muon"); + py::class_(m, "Allocator") + .def(py::init<>()); + + py::class_(m, "HypersT") + .def_readwrite("horizon", &HypersT::horizon) + .def_readwrite("total_agents", &HypersT::total_agents) + .def_readwrite("num_buffers", &HypersT::num_buffers) + .def_readwrite("num_atns", &HypersT::num_atns) + .def_readwrite("hidden_size", &HypersT::hidden_size) + + .def_readwrite("replay_ratio", &HypersT::replay_ratio) + .def_readwrite("num_layers", &HypersT::num_layers) + .def_readwrite("lr", &HypersT::lr) + .def_readwrite("min_lr_ratio", &HypersT::min_lr_ratio) + .def_readwrite("anneal_lr", &HypersT::anneal_lr) + .def_readwrite("beta1", &HypersT::beta1) + .def_readwrite("beta2", &HypersT::beta2) + .def_readwrite("eps", &HypersT::eps) + .def_readwrite("total_timesteps", &HypersT::total_timesteps) + .def_readwrite("max_grad_norm", &HypersT::max_grad_norm) + .def_readwrite("clip_coef", &HypersT::clip_coef) + .def_readwrite("vf_clip_coef", &HypersT::vf_clip_coef) + .def_readwrite("vf_coef", &HypersT::vf_coef) + .def_readwrite("ent_coef", &HypersT::ent_coef) + .def_readwrite("min_ent_coef_ratio", &HypersT::min_ent_coef_ratio) + .def_readwrite("anneal_ent_coef", &HypersT::anneal_ent_coef) + .def_readwrite("gamma", &HypersT::gamma) + .def_readwrite("gae_lambda", &HypersT::gae_lambda) + .def_readwrite("vtrace_rho_clip", &HypersT::vtrace_rho_clip) + .def_readwrite("vtrace_c_clip", &HypersT::vtrace_c_clip) + .def_readwrite("prio_alpha", &HypersT::prio_alpha) + .def_readwrite("prio_beta0", &HypersT::prio_beta0) + .def_readwrite("cudagraphs", &HypersT::cudagraphs) + .def_readwrite("profile", &HypersT::profile) + .def_readwrite("rank", &HypersT::rank) + .def_readwrite("world_size", &HypersT::world_size) + .def_readwrite("gpu_id", &HypersT::gpu_id) + .def_readwrite("nccl_id", &HypersT::nccl_id); + + py::class_(m, "PrecisionTensor") + .def("__repr__", [](const PrecisionTensor& t) { return std::string(puf_repr(&t)); }) + .def("ndim", [](const PrecisionTensor& t) { return ndim(t.shape); }) + .def("numel", [](const PrecisionTensor& t) { return numel(t.shape); }); + py::class_(m, "FloatTensor") + .def("__repr__", [](const FloatTensor& t) { return std::string(puf_repr(&t)); }) + .def("ndim", [](const FloatTensor& t) { return ndim(t.shape); }) + .def("numel", [](const FloatTensor& t) { return numel(t.shape); }); + + py::class_(m, "RolloutBuf") + .def_readwrite("observations", &RolloutBuf::observations) + .def_readwrite("actions", &RolloutBuf::actions) + .def_readwrite("values", &RolloutBuf::values) + .def_readwrite("logprobs", &RolloutBuf::logprobs) + .def_readwrite("rewards", &RolloutBuf::rewards) + .def_readwrite("terminals", &RolloutBuf::terminals) + .def_readwrite("ratio", &RolloutBuf::ratio) + .def_readwrite("importance", &RolloutBuf::importance); + + m.def("uptime", [](py::object pufferl_obj) -> double { + PuffeRL& pufferl = pufferl_obj.cast(); + double now = wall_clock(); + return now - pufferl.start_time; + }); + m.def("puff_advantage", &py_puff_advantage); + m.def("create_vec", &create_vec, py::arg("args"), py::arg("gpu") = 1); + py::class_>(m, "VecEnv") + .def_readonly("total_agents", &VecEnv::total_agents) + .def_readonly("obs_size", &VecEnv::obs_size) + .def_readonly("num_atns", &VecEnv::num_atns) + .def_readonly("act_sizes", &VecEnv::act_sizes) + .def_readonly("obs_dtype", &VecEnv::obs_dtype) + .def_readonly("obs_elem_size", &VecEnv::obs_elem_size) + .def_readonly("gpu", &VecEnv::gpu) + // GPU buffer pointers — wrap with torch.from_blob(..., device='cuda') + .def_property_readonly("gpu_obs_ptr", [](VecEnv& ve) { return (long long)ve.vec->gpu_observations; }) + .def_property_readonly("gpu_rewards_ptr", [](VecEnv& ve) { return (long long)ve.vec->gpu_rewards; }) + .def_property_readonly("gpu_terminals_ptr", [](VecEnv& ve) { return (long long)ve.vec->gpu_terminals; }) + // CPU buffer pointers (same as gpu_ in CPU mode since they alias) + .def_property_readonly("obs_ptr", [](VecEnv& ve) { return (long long)ve.vec->observations; }) + .def_property_readonly("rewards_ptr", [](VecEnv& ve) { return (long long)ve.vec->rewards; }) + .def_property_readonly("terminals_ptr", [](VecEnv& ve) { return (long long)ve.vec->terminals; }) + .def("reset", &vec_reset) + .def("gpu_step", &gpu_vec_step_py) + .def("cpu_step", &cpu_vec_step_py) + .def("render", [](VecEnv& ve, int env_id) { static_vec_render(ve.vec, env_id); }) + .def("log", &vec_log) + .def("close", &vec_close); + + m.def("create_pufferl", &create_pufferl); + py::class_>(m, "PuffeRL") + .def_readwrite("policy", &PuffeRL::policy) + .def_readwrite("muon", &PuffeRL::muon) + .def_readwrite("hypers", &PuffeRL::hypers) + .def_readwrite("rollouts", &PuffeRL::rollouts) + .def_readonly("epoch", &PuffeRL::epoch) + .def_readonly("global_step", &PuffeRL::global_step) + .def_readonly("last_log_time", &PuffeRL::last_log_time) + .def("num_params", [](PuffeRL& self) -> int64_t { + return numel(self.master_weights.shape); + }); +} diff --git a/src-hip/cuda_shim.cpp b/src-hip/cuda_shim.cpp new file mode 100644 index 0000000000..8f27002933 --- /dev/null +++ b/src-hip/cuda_shim.cpp @@ -0,0 +1,18 @@ +// cuda->HIP forwarding shims for the C env libraries on AMD. +#include +typedef int cudaError_t; +extern "C" { +cudaError_t cudaHostAlloc(void** p, size_t s, unsigned int f) { return hipHostAlloc(p, s, f); } +cudaError_t cudaMalloc(void** p, size_t s) { return hipMalloc(p, s); } +cudaError_t cudaMemcpy(void* d, const void* s, size_t n, int k) { return hipMemcpy(d, s, n, (hipMemcpyKind)k); } +cudaError_t cudaMemcpyAsync(void* d, const void* s, size_t n, int k, void* st) { return hipMemcpyAsync(d, s, n, (hipMemcpyKind)k, (hipStream_t)st); } +cudaError_t cudaMemset(void* d, int v, size_t n) { return hipMemset(d, v, n); } +cudaError_t cudaFree(void* p) { return hipFree(p); } +cudaError_t cudaFreeHost(void* p) { return hipHostFree(p); } +cudaError_t cudaSetDevice(int i) { return hipSetDevice(i); } +cudaError_t cudaDeviceSynchronize(void) { return hipDeviceSynchronize(); } +cudaError_t cudaStreamSynchronize(void* s) { return hipStreamSynchronize((hipStream_t)s); } +cudaError_t cudaStreamCreateWithFlags(void** s, unsigned int f) { return hipStreamCreateWithFlags((hipStream_t*)s, f); } +cudaError_t cudaStreamQuery(void* s) { return hipStreamQuery((hipStream_t)s); } +const char* cudaGetErrorString(cudaError_t e) { return hipGetErrorString((hipError_t)e); } +} diff --git a/src-hip/cudnn_conv2d.hip.cpp b/src-hip/cudnn_conv2d.hip.cpp new file mode 100644 index 0000000000..1598ce89fa --- /dev/null +++ b/src-hip/cudnn_conv2d.hip.cpp @@ -0,0 +1,231 @@ +// MIOpen Conv2d: forward/backward with fused bias+activation. +// AMD counterpart of src/cudnn_conv2d.cu (which keeps the NVIDIA/cuDNN +// implementation unchanged). Same struct and function interface, so +// ocean.hip.cpp can include this file where ocean.cu includes the cuDNN +// original. Included by ocean.hip.cpp (training). + +#ifndef CUDNN_CONV2D_CU +#define CUDNN_CONV2D_CU +#ifndef CUDNN_CONV2D_CU +#define CUDNN_CONV2D_CU + +#include +#include +#include +#include + +#include "kernels.hip.cpp" + +#ifndef CHECK_MIOPEN +#define CHECK_MIOPEN(call) do { \ + miopenStatus_t e = call; \ + if (e != miopenStatusSuccess) { \ + fprintf(stderr, "MIOpen %s:%d: %s\n", __FILE__, __LINE__, miopenGetErrorString(e)); exit(1); \ + } \ +} while(0) +#endif + +static inline int div_ceil(int a, int b) { return (a + b - 1) / b; } + +static miopenHandle_t get_cudnn_handle() { + static miopenHandle_t h = nullptr; + if (!h) CHECK_MIOPEN(miopenCreate(&h)); + return h; +} + +__global__ void relu_inplace_kernel(precision_t* data, int n) { + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < n && data[i] < 0) data[i] = precision_t(0); +} + +// ---- ConvWeights: params + batch-independent MIOpen state ---- + +struct ConvWeights { + PrecisionTensor w, b; // w: (OC, IC*K*K), b: (OC) + int IC, OC, K, S, IH, IW, OH, OW; + bool relu; + miopenDataType_t dtype; + miopenTensorDescriptor_t cudnn_bias; + miopenTensorDescriptor_t cudnn_filt; + miopenConvolutionDescriptor_t cudnn_conv; + miopenActivationDescriptor_t cudnn_act; + bool cudnn_ready; +}; + +// ---- ConvActivations: per-batch-size buffers + descriptors ---- + +struct ConvActivations { + PrecisionTensor out, grad, saved_input; + PrecisionTensor wgrad, bgrad; + // Per-batch-size MIOpen state + miopenTensorDescriptor_t cudnn_in, cudnn_out; + miopenConvFwdAlgorithm_t fwd_algo; + miopenConvBwdDataAlgorithm_t bwd_data_algo; + miopenConvBwdWeightsAlgorithm_t bwd_filt_algo; + size_t fwd_ws_bytes, bwd_data_ws_bytes, bwd_filt_ws_bytes; + void* fwd_ws; void* bwd_data_ws; void* bwd_filt_ws; + bool cudnn_setup; +}; + +static void conv_init(ConvWeights* cw, int IC, int OC, int K, int S, int IH, int IW, bool relu) { + cw->IC = IC; cw->OC = OC; cw->K = K; cw->S = S; cw->IH = IH; cw->IW = IW; + cw->OH = (IH - K) / S + 1; cw->OW = (IW - K) / S + 1; + cw->relu = relu; cw->cudnn_ready = false; +} + +// Create batch-independent descriptors (once) +static void conv_setup_shared(ConvWeights* cw, miopenDataType_t dt) { + if (cw->cudnn_ready) return; + assert(dt == miopenFloat && "MIOpen conv path supports float only"); + cw->dtype = dt; + CHECK_MIOPEN(miopenCreateTensorDescriptor(&cw->cudnn_filt)); + CHECK_MIOPEN(miopenSet4dTensorDescriptor(cw->cudnn_filt, dt, cw->OC, cw->IC, cw->K, cw->K)); + CHECK_MIOPEN(miopenCreateConvolutionDescriptor(&cw->cudnn_conv)); + CHECK_MIOPEN(miopenInitConvolutionDescriptor(cw->cudnn_conv, miopenConvolution, + 0, 0, cw->S, cw->S, 1, 1)); + CHECK_MIOPEN(miopenCreateTensorDescriptor(&cw->cudnn_bias)); + CHECK_MIOPEN(miopenSet4dTensorDescriptor(cw->cudnn_bias, dt, 1, cw->OC, 1, 1)); + CHECK_MIOPEN(miopenCreateActivationDescriptor(&cw->cudnn_act)); + CHECK_MIOPEN(miopenSetActivationDescriptor(cw->cudnn_act, + cw->relu ? miopenActivationRELU : miopenActivationPASTHRU, 0.0, 0.0, 0.0)); + cw->cudnn_ready = true; +} + +// Setup per-activation-set MIOpen state: batch-dependent descriptors + algo search + workspace +static void conv_setup_activations(ConvWeights* cw, ConvActivations* ca, int B, miopenDataType_t dt) { + conv_setup_shared(cw, dt); + miopenHandle_t h = get_cudnn_handle(); + + CHECK_MIOPEN(miopenCreateTensorDescriptor(&ca->cudnn_in)); + CHECK_MIOPEN(miopenSet4dTensorDescriptor(ca->cudnn_in, dt, B, cw->IC, cw->IH, cw->IW)); + CHECK_MIOPEN(miopenCreateTensorDescriptor(&ca->cudnn_out)); + CHECK_MIOPEN(miopenSet4dTensorDescriptor(ca->cudnn_out, dt, B, cw->OC, cw->OH, cw->OW)); + + int returned; + size_t ws_bytes = 0; + + miopenConvAlgoPerf_t fp; + CHECK_MIOPEN(miopenFindConvolutionForwardAlgorithm(h, + ca->cudnn_in, nullptr, cw->cudnn_filt, cw->w.data, cw->cudnn_conv, + ca->cudnn_out, nullptr, 1, &returned, &fp, + nullptr, 0, false)); + ca->fwd_algo = fp.fwd_algo; + CHECK_MIOPEN(miopenConvolutionForwardGetWorkSpaceSize(h, + cw->cudnn_filt, ca->cudnn_in, cw->cudnn_conv, ca->cudnn_out, + &ws_bytes)); + ca->fwd_ws_bytes = ws_bytes; + ca->fwd_ws = NULL; if (ca->fwd_ws_bytes > 0) hipMalloc(&ca->fwd_ws, ca->fwd_ws_bytes); + + miopenConvAlgoPerf_t ffp; + CHECK_MIOPEN(miopenFindConvolutionBackwardWeightsAlgorithm(h, + ca->cudnn_out, nullptr, ca->cudnn_in, nullptr, cw->cudnn_conv, + cw->cudnn_filt, ca->wgrad.data, 1, &returned, &ffp, + nullptr, 0, false)); + ca->bwd_filt_algo = ffp.bwd_weights_algo; + CHECK_MIOPEN(miopenConvolutionBackwardWeightsGetWorkSpaceSize(h, + ca->cudnn_out, ca->cudnn_in, cw->cudnn_conv, cw->cudnn_filt, + &ws_bytes)); + ca->bwd_filt_ws_bytes = ws_bytes; + ca->bwd_filt_ws = NULL; if (ca->bwd_filt_ws_bytes > 0) hipMalloc(&ca->bwd_filt_ws, ca->bwd_filt_ws_bytes); + + miopenConvAlgoPerf_t dp; + CHECK_MIOPEN(miopenFindConvolutionBackwardDataAlgorithm(h, + cw->cudnn_filt, nullptr, ca->cudnn_out, nullptr, cw->cudnn_conv, + ca->cudnn_in, nullptr, 1, &returned, &dp, + nullptr, 0, false)); + ca->bwd_data_algo = dp.bwd_data_algo; + CHECK_MIOPEN(miopenConvolutionBackwardDataGetWorkSpaceSize(h, + ca->cudnn_out, cw->cudnn_filt, cw->cudnn_conv, ca->cudnn_in, + &ws_bytes)); + ca->bwd_data_ws_bytes = ws_bytes; + ca->bwd_data_ws = NULL; if (ca->bwd_data_ws_bytes > 0) hipMalloc(&ca->bwd_data_ws, ca->bwd_data_ws_bytes); + + ca->cudnn_setup = true; +} + +// Legacy single-setup API (for tests) +static void conv_setup(ConvWeights* cw, int B, miopenDataType_t dt) { + (void)B; + conv_setup_shared(cw, dt); +} + +static void conv_reg_params(ConvWeights* cw, Allocator* alloc) { + cw->w = {.shape = {cw->OC, cw->IC * cw->K * cw->K}}; + cw->b = {.shape = {cw->OC}}; + alloc_register(alloc,&cw->w); alloc_register(alloc,&cw->b); +} + +static void conv_reg_train(ConvWeights* cw, ConvActivations* ca, Allocator* acts, Allocator* grads, int B, miopenDataType_t dt) { + ca->out = {.shape = {B * cw->OC * cw->OH * cw->OW}}; + ca->grad = {.shape = {B * cw->OC * cw->OH * cw->OW}}; + ca->saved_input = {.shape = {B * cw->IC * cw->IH * cw->IW}}; + ca->wgrad = {.shape = {cw->OC, cw->IC * cw->K * cw->K}}; + ca->bgrad = {.shape = {cw->OC}}; + alloc_register(acts,&ca->out); alloc_register(acts,&ca->grad); alloc_register(acts,&ca->saved_input); + alloc_register(grads,&ca->wgrad); alloc_register(grads,&ca->bgrad); + conv_setup_activations(cw, ca, B, dt); +} + +static void conv_reg_rollout(ConvWeights* cw, ConvActivations* ca, Allocator* alloc, int B, miopenDataType_t dt) { + ca->out = {.shape = {B * cw->OC * cw->OH * cw->OW}}; + ca->cudnn_setup = false; + alloc_register(alloc,&ca->out); + conv_setup_activations(cw, ca, B, dt); +} + +static void conv_init_weights(ConvWeights* cw, uint64_t* seed, hipStream_t stream) { + PrecisionTensor wt = {.data = cw->w.data, .shape = {cw->OC, cw->IC * cw->K * cw->K}}; + puf_kaiming_init(&wt, 1.0f, (*seed)++, stream); + hipMemsetAsync(cw->b.data, 0, numel(cw->b.shape) * sizeof(precision_t), stream); +} + +// ---- Forward / Backward ---- + +// Fused conv + bias + activation. All NCHW. Saves input for backward. +static void conv_forward(ConvWeights* cw, ConvActivations* ca, void* input, int B, hipStream_t stream) { + miopenHandle_t h = get_cudnn_handle(); + CHECK_MIOPEN(miopenSetStream(h, stream)); + float alpha = 1.0f, beta = 0.0f; + if (ca->saved_input.data) { + hipMemcpyAsync(ca->saved_input.data, input, + (int64_t)B * cw->IC * cw->IH * cw->IW * sizeof(precision_t), hipMemcpyDeviceToDevice, stream); + } + CHECK_MIOPEN(miopenConvolutionForward(h, + &alpha, ca->cudnn_in, input, cw->cudnn_filt, cw->w.data, + cw->cudnn_conv, ca->fwd_algo, + &beta, ca->cudnn_out, ca->out.data, + ca->fwd_ws, ca->fwd_ws_bytes)); + CHECK_MIOPEN(miopenConvolutionForwardBias(h, + &alpha, cw->cudnn_bias, cw->b.data, + &beta, ca->cudnn_out, ca->out.data)); + if (cw->relu) { + int n = B * cw->OC * cw->OH * cw->OW; + relu_inplace_kernel<<>>( + (precision_t*)ca->out.data, n); + } +} + +// Backward: upstream grad in ca->grad, relu mask in ca->out. +// Caller must apply relu backward and bias grad (dtype-specific kernels). +// This does MIOpen filter grad + optional data grad. +static void conv_backward(ConvWeights* cw, ConvActivations* ca, void* input_grad, int B, hipStream_t stream) { + (void)B; + miopenHandle_t h = get_cudnn_handle(); + CHECK_MIOPEN(miopenSetStream(h, stream)); + float alpha = 1.0f, beta = 0.0f; + + CHECK_MIOPEN(miopenConvolutionBackwardWeights(h, + &alpha, ca->cudnn_out, ca->grad.data, ca->cudnn_in, ca->saved_input.data, + cw->cudnn_conv, ca->bwd_filt_algo, &beta, cw->cudnn_filt, ca->wgrad.data, + ca->bwd_filt_ws, ca->bwd_filt_ws_bytes)); + + if (input_grad) { + CHECK_MIOPEN(miopenConvolutionBackwardData(h, + &alpha, ca->cudnn_out, ca->grad.data, cw->cudnn_filt, cw->w.data, + cw->cudnn_conv, ca->bwd_data_algo, + &beta, ca->cudnn_in, input_grad, + ca->bwd_data_ws, ca->bwd_data_ws_bytes)); + } +} + +#endif // CUDNN_CONV2D_CU diff --git a/src-hip/kernels.hip.cpp b/src-hip/kernels.hip.cpp new file mode 100644 index 0000000000..3db984c658 --- /dev/null +++ b/src-hip/kernels.hip.cpp @@ -0,0 +1,470 @@ +#include "hip/hip_runtime.h" +#ifndef PUFFERLIB_KERNELS_CU +#define PUFFERLIB_KERNELS_CU + +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +#ifdef PRECISION_FLOAT +typedef float precision_t; +constexpr bool USE_BF16 = false; +constexpr int PRECISION_SIZE = 4; +static constexpr hipDataType CUBLAS_PRECISION = HIP_R_32F; +static constexpr hipblasComputeType_t CUBLAS_COMPUTE_PRECISION = HIPBLAS_COMPUTE_32F; +#define NCCL_PRECISION ncclFloat +#define to_float(x) (x) +#define from_float(x) (x) +#else +typedef __hip_bfloat16 precision_t; +constexpr bool USE_BF16 = true; +constexpr int PRECISION_SIZE = 2; +static constexpr hipDataType CUBLAS_PRECISION = HIP_R_16BF; +static constexpr hipblasComputeType_t CUBLAS_COMPUTE_PRECISION = HIPBLAS_COMPUTE_32F; +#define NCCL_PRECISION ncclBfloat16 +#define to_float(x) __bfloat162float(x) +#define from_float(x) __float2bfloat16(x) +#endif + +#include "tensor.h" + +__host__ __device__ inline int ndim(const int64_t* shape) { + int n = 0; while (n < PUF_MAX_DIMS && shape[n] != 0) n++; return n; +} + +__host__ __device__ inline int64_t numel(const int64_t* shape) { + int64_t n = 1; for (int i = 0; i < PUF_MAX_DIMS && shape[i] != 0; i++) n *= shape[i]; return n; +} + +inline int64_t batch_size(const int64_t* shape) { + int n = ndim(shape); + int64_t b = 1; + for (int i = 0; i < n - 2; i++) b *= shape[i]; + return b; +} + +inline const char* _puf_repr_impl(const char* name, const char* dtype, + const int64_t* shape, int nd, int64_t ne, bool empty) { + static thread_local char buf[256]; + if (empty) { snprintf(buf, sizeof(buf), "%s(empty)", name); return buf; } + int pos = snprintf(buf, sizeof(buf), "%s(%s, [", name, dtype); + for (int i = 0; i < nd && pos < (int)sizeof(buf) - 32; i++) + pos += snprintf(buf + pos, sizeof(buf) - pos, "%s%lld", i ? ", " : "", (long long)shape[i]); + snprintf(buf + pos, sizeof(buf) - pos, "], %lld elems)", (long long)ne); + return buf; +} + +inline const char* puf_repr(const PrecisionTensor* t) { + return _puf_repr_impl("PrecisionTensor", USE_BF16 ? "bf16" : "f32", + t->shape, ndim(t->shape), numel(t->shape), !t->data); +} + +inline const char* puf_repr(const FloatTensor* t) { + return _puf_repr_impl("FloatTensor", "f32", + t->shape, ndim(t->shape), numel(t->shape), !t->data); +} + +#ifndef HIP_INF_F +#define HIP_INF_F __int_as_float(0x7f800000) +#endif + +#define PPO_THREADS 256 +#define SELECT_COPY_THREADS 256 +#define MAX_ATN_HEADS 16 + + +#define BLOCK_SIZE 256 +inline int grid_size(int N) { + return (N + BLOCK_SIZE - 1) / BLOCK_SIZE; +} + +#define SEQ_SIZE 256 +inline int seq_size(int N) { + return (N + SEQ_SIZE - 1) / SEQ_SIZE; +} + +#define SOFTPLUS_BETA 1.0f +#define SOFTPLUS_THRESHOLD 20.0f +__device__ __forceinline__ float softplus_fwd(float x) { + float x_scaled = x * SOFTPLUS_BETA; + return (x_scaled > SOFTPLUS_THRESHOLD) ? x : log1pf(expf(x_scaled)) / SOFTPLUS_BETA; +} + +__device__ __forceinline__ float softplus_bwd(float grad_output, float x) { + float beta_x = SOFTPLUS_BETA * x; + if (beta_x > SOFTPLUS_THRESHOLD) { + return grad_output; + } + float exp_beta_x = expf(beta_x); + return grad_output * (exp_beta_x / (1.0f + exp_beta_x)); +} + +__device__ __forceinline__ float relu(float x) { + return fmaxf(0.0f, x); +} + +__device__ __forceinline__ float relu_backward(float x, float grad_output) { + return (x > 0.0f) ? grad_output : 0.0f; +} + +__device__ __forceinline__ float sigmoid(float x) { + float z = expf(-fabsf(x)); + return x >= 0.0f ? 1.0f / (1.0f + z) : z / (1.0f + z); +} + +__device__ __forceinline__ float sigmoid_backward(float x, float grad_output) { + float sig = sigmoid(x); + return grad_output * sig * (1.0f - sig); +} + +__device__ __inline__ float fast_tanh(float x) { + float v1 = fminf(fmaxf(x, -9.0f), 9.0f); + float v2 = v1 * v1; + float p = v2 * -2.76076847742355e-16f + 2.00018790482477e-13f; + p = v2 * p + -8.60467152213735e-11f; + p = v2 * p + 5.12229709037114e-08f; + p = v2 * p + 1.48572235717979e-05f; + p = v2 * p + 6.37261928875436e-04f; + p = v2 * p + 4.89352455891786e-03f; + p = v1 * p; + float q = v2 * 1.19825839466702e-06f + 1.18534705686654e-04f; + q = v2 * q + 2.26843463243900e-03f; + q = v2 * q + 4.89352518554385e-03f; + return p / q; +} + +__device__ __inline__ float fast_sigmoid(float x) { + return fminf(1.0f, fmaxf(0.0f, (fast_tanh(x * 0.5f) + 1.0f) * 0.5f)); +} + +__device__ __forceinline__ float lerp(float a, float b, float w) { + float diff = b - a; + return (fabsf(w) < 0.5f) ? a + w * diff : b - diff * (1.0f - w); +} + +__device__ __forceinline__ float logaddexp(float a, float b) { + float m = fmaxf(a, b), diff = fminf(a, b) - m; + return (diff < -88.0f) ? m : m + log1pf(__expf(diff)); +} + +//TODO: Speed up. The previous version was misaligned. +__device__ __forceinline__ void copy_bytes( + const char* __restrict__ src, char* __restrict__ dst, + int src_row, int dst_row, int row_bytes) { + const char* s = src + (int64_t)src_row * row_bytes; + char* d = dst + (int64_t)dst_row * row_bytes; + for (int i = threadIdx.x; i < row_bytes; i += blockDim.x) { + d[i] = s[i]; + } +} + +/* +__device__ __forceinline__ void copy_bytes(const char* __restrict__ src, + char* __restrict__ dst, int src_row, int dst_row, int row_bytes) { + const int* soffset = (const int*)(src + (int64_t)src_row * row_bytes); + int* doffset = (int*)(dst + (int64_t)dst_row * row_bytes); + for (int i = threadIdx.x; i < row_bytes / 4; i += blockDim.x) { + doffset[i] = soffset[i]; + } +} +*/ + +// Transpose dims 0,1: [A, B, C] -> [B, A, C]. For 2D, pass C=1. +__global__ void transpose_102(precision_t* __restrict__ dst, + const precision_t* __restrict__ src, int A, int B, int C) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = A * B * C; + if (idx >= total) { + return; + } + int a = idx / (B * C), rem = idx % (B * C), b = rem / C, c = rem % C; + dst[b * A * C + a * C + c] = src[idx]; +} + +__global__ void fill_precision_kernel(precision_t* __restrict__ dst, precision_t val, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = val; + } +} + +__global__ void clamp_precision_kernel(precision_t* __restrict__ dst, float lo, float hi, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + float v = to_float(dst[idx]); + dst[idx] = from_float(fminf(fmaxf(v, lo), hi)); + } +} + +__global__ void add_kernel(float* __restrict__ dst, const precision_t* __restrict__ src, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] += to_float(src[idx]); + } +} + +#ifndef PRECISION_FLOAT +__global__ void add_kernel(precision_t* __restrict__ dst, const precision_t* __restrict__ src, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = from_float(to_float(dst[idx]) + to_float(src[idx])); + } +} +#endif + +// merge shape[dim] into shape[dim+1] --- +inline PrecisionTensor* puf_squeeze(PrecisionTensor* t, int dim) { + int n = ndim(t->shape); + t->shape[dim + 1] *= t->shape[dim]; + for (int i = dim; i < n - 1; i++) t->shape[i] = t->shape[i + 1]; + t->shape[n - 1] = 0; + return t; +} +inline FloatTensor* puf_squeeze(FloatTensor* t, int dim) { + int n = ndim(t->shape); + t->shape[dim + 1] *= t->shape[dim]; + for (int i = dim; i < n - 1; i++) t->shape[i] = t->shape[i + 1]; + t->shape[n - 1] = 0; + return t; +} + +// split shape[dim] into {d0, d1} --- +inline PrecisionTensor* puf_unsqueeze(PrecisionTensor* t, int dim, int64_t d0, int64_t d1) { + assert(d0 * d1 == t->shape[dim] && "puf_unsqueeze: d0 * d1 must equal shape[dim]"); + int n = ndim(t->shape); + for (int i = n; i > dim; i--) t->shape[i] = t->shape[i - 1]; + t->shape[dim] = d0; + t->shape[dim + 1] = d1; + return t; +} + +// Dense row-major GEMM: C(M,N) = alpha * op_a(A) @ op_b(B) + beta * C +// Strides derived from M, N, K assuming tightly packed row-major storage. +static const size_t CUBLAS_WS_BYTES = 32 * 1024 * 1024; + +static hipblasHandle_t cublas_get_handle() { + static thread_local hipblasHandle_t handle = nullptr; + if (!handle) { + hipblasCreate(&handle); + void* ws = nullptr; + hipMalloc(&ws, CUBLAS_WS_BYTES); + hipblasSetWorkspace(handle, ws, CUBLAS_WS_BYTES); + } + return handle; +} + +static inline void cublasGemmExDense( + hipblasOperation_t op_a, hipblasOperation_t op_b, + int M, int N, int K, void* A, void* B, void* C, + hipStream_t stream, float alpha = 1.0f, float beta = 0.0f) { + int lda = (op_a == HIPBLAS_OP_N) ? K : M; + int ldb = (op_b == HIPBLAS_OP_N) ? N : K; + + hipblasHandle_t handle = cublas_get_handle(); + hipblasSetStream(handle, stream); + hipblasGemmEx(handle, op_b, op_a, N, M, K, &alpha, + B, CUBLAS_PRECISION, ldb, A, CUBLAS_PRECISION, lda, &beta, + C, CUBLAS_PRECISION, N, CUBLAS_COMPUTE_PRECISION, HIPBLAS_GEMM_DEFAULT); +} + +// out(...,N) = a(...,K) @ b(N,K)^T — leading dims folded into M +void puf_mm(PrecisionTensor* a, PrecisionTensor* b, PrecisionTensor* out, hipStream_t stream) { + int M = batch_size(a->shape) * a->shape[ndim(a->shape)-2]; + int K = a->shape[ndim(a->shape)-1]; + int N = b->shape[ndim(b->shape)-2]; + cublasGemmExDense(HIPBLAS_OP_N, HIPBLAS_OP_T, M, N, K, + a->data, b->data, out->data, stream); +} + +// out(M,N) = a(...,M)^T @ b(...,N) — leading dims folded into K +void puf_mm_tn(PrecisionTensor* a, PrecisionTensor* b, PrecisionTensor* out, hipStream_t stream) { + int M = a->shape[ndim(a->shape)-1]; + int K = batch_size(a->shape) * a->shape[ndim(a->shape)-2]; + int N = b->shape[ndim(b->shape)-1]; + cublasGemmExDense(HIPBLAS_OP_T, HIPBLAS_OP_N, M, N, K, + a->data, b->data, out->data, stream); +} + +// out(...,N) = a(...,K) @ b(K,N) — leading dims folded into M +void puf_mm_nn(PrecisionTensor* a, PrecisionTensor* b, PrecisionTensor* out, hipStream_t stream) { + int M = batch_size(a->shape) * a->shape[ndim(a->shape)-2]; + int K = a->shape[ndim(a->shape)-1]; + int N = b->shape[ndim(b->shape)-1]; + cublasGemmExDense(HIPBLAS_OP_N, HIPBLAS_OP_N, M, N, K, + a->data, b->data, out->data, stream); +} + +static void puf_addmm_nn(PrecisionTensor* a, PrecisionTensor* b, PrecisionTensor* out, + float alpha, float beta, hipStream_t stream) { + int M = batch_size(a->shape) * a->shape[ndim(a->shape)-2]; + int K = a->shape[ndim(a->shape)-1]; + int N = b->shape[ndim(b->shape)-1]; + cublasGemmExDense(HIPBLAS_OP_N, HIPBLAS_OP_N, M, N, K, + a->data, b->data, out->data, stream, alpha, beta); +} + +__global__ void cast(precision_t* __restrict__ dst, + const float* __restrict__ src, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = from_float(src[idx]); + } +} + +#ifndef PRECISION_FLOAT +inline void cast_dispatch(precision_t* dst, const precision_t* src, int n, hipStream_t stream) { + hipMemcpyAsync(dst, src, n * sizeof(precision_t), hipMemcpyDeviceToDevice, stream); +} +#endif +inline void cast_dispatch(precision_t* dst, const float* src, int n, hipStream_t stream) { + cast<<>>(dst, src, n); +} + +#ifndef PRECISION_FLOAT +__global__ void cast(float* __restrict__ dst, + const precision_t* __restrict__ src, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = to_float(src[idx]); + } +} +#endif + +__global__ void cast(precision_t* __restrict__ dst, + const unsigned char* __restrict__ src, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = from_float((float)src[idx]); + } +} + +__global__ void cast(unsigned char* __restrict__ dst, + const precision_t* __restrict__ src, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = to_float(src[idx]); + } +} + +inline void cast_dispatch(precision_t* dst, const unsigned char* src, int n, hipStream_t stream) { + cast<<>>(dst, src, n); +} + +void puf_copy(PrecisionTensor* dst, const PrecisionTensor* src, hipStream_t stream) { + assert(numel(dst->shape) == numel(src->shape) && "puf_copy: size mismatch"); + hipMemcpyAsync(dst->data, src->data, numel(dst->shape) * sizeof(precision_t), hipMemcpyDeviceToDevice, stream); +} + +void puf_zero(PrecisionTensor* dst, hipStream_t stream) { + hipMemsetAsync(dst->data, 0, numel(dst->shape) * sizeof(precision_t), stream); +} + +void puf_zero(FloatTensor* dst, hipStream_t stream) { + hipMemsetAsync(dst->data, 0, numel(dst->shape) * sizeof(float), stream); +} + +__global__ void uniform_scale_kernel(float* data, float bound, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) data[idx] = data[idx] * 2.0f * bound - bound; +} + +// Uniform(-1/sqrt(fan_in), 1/sqrt(fan_in)) +void puf_kaiming_init(PrecisionTensor* dst, float gain, ulong seed, hipStream_t stream) { + assert(ndim(dst->shape) == 2); + long rows = dst->shape[0], cols = dst->shape[1]; + assert(rows > 0 && cols > 0); + long n = rows * cols; + float bound = gain / std::sqrt((float)cols); + float* buf; + hipMalloc(&buf, n * sizeof(float)); + hiprandGenerator_t gen; + hiprandCreateGenerator(&gen, HIPRAND_RNG_PSEUDO_DEFAULT); + hiprandSetPseudoRandomGeneratorSeed(gen, seed); + hiprandGenerateUniform(gen, buf, n); + hiprandDestroyGenerator(gen); + uniform_scale_kernel<<>>(buf, bound, n); + cast<<>>(dst->data, buf, n); + hipFree(buf); +} + +// Normal(0, std). Used for embeddings +void puf_normal_init(PrecisionTensor* dst, float std, ulong seed, hipStream_t stream) { + long n = numel(dst->shape); + assert(n > 0); + long rand_count = (n % 2 == 0) ? n : n + 1; + float* buf; + hipMalloc(&buf, rand_count * sizeof(float)); + hiprandGenerator_t gen; + hiprandCreateGenerator(&gen, HIPRAND_RNG_PSEUDO_DEFAULT); + hiprandSetPseudoRandomGeneratorSeed(gen, seed); + hiprandGenerateNormal(gen, buf, rand_count, 0.0f, std); + hiprandDestroyGenerator(gen); + cast<<>>(dst->data, buf, n); + hipFree(buf); +} + +struct AllocEntry { + void** data_ptr; // address of the tensor's data field + int64_t* shape; // pointer to the tensor's shape array + int elem_size; // sizeof element type +}; + +struct Allocator { + AllocEntry* regs = nullptr; + int num_regs = 0; + void* mem = nullptr; + long total_elems = 0; + long total_bytes = 0; +}; + +static void alloc_register_impl(Allocator* alloc, void** data_ptr, int64_t* shape, int elem_size) { + alloc->regs = (AllocEntry*)realloc(alloc->regs, (alloc->num_regs + 1) * sizeof(AllocEntry)); + alloc->regs[alloc->num_regs++] = {data_ptr, shape, elem_size}; + int64_t n = numel(shape); + alloc->total_elems += n; + alloc->total_bytes = (alloc->total_bytes + 15) & ~15; + alloc->total_bytes += n * elem_size; +} +void alloc_register(Allocator* a, PrecisionTensor* t) { + alloc_register_impl(a, (void**)&t->data, t->shape, sizeof(precision_t)); +} +void alloc_register(Allocator* a, FloatTensor* t) { + alloc_register_impl(a, (void**)&t->data, t->shape, sizeof(float)); +} +void alloc_register(Allocator* a, LongTensor* t) { + alloc_register_impl(a, (void**)&t->data, t->shape, sizeof(long)); +} +void alloc_register(Allocator* a, IntTensor* t) { + alloc_register_impl(a, (void**)&t->data, t->shape, sizeof(int)); +} + +hipError_t alloc_create(Allocator* alloc) { + if (alloc->total_bytes == 0) return hipSuccess; + hipError_t err = hipMalloc(&alloc->mem, alloc->total_bytes); + if (err != hipSuccess) return err; + hipMemset(alloc->mem, 0, alloc->total_bytes); + long offset = 0; + for (int i = 0; i < alloc->num_regs; i++) { + offset = (offset + 15) & ~15; + *alloc->regs[i].data_ptr = (char*)alloc->mem + offset; + offset += numel(alloc->regs[i].shape) * alloc->regs[i].elem_size; + } + return hipSuccess; +} + +void alloc_free(Allocator* alloc) { + if (alloc->mem) { hipFree(alloc->mem); alloc->mem = nullptr; } + if (alloc->regs) { free(alloc->regs); alloc->regs = nullptr; } + alloc->num_regs = 0; + alloc->total_elems = 0; + alloc->total_bytes = 0; +} + +#endif // PUFFERLIB_KERNELS_CU diff --git a/src-hip/models.hip.cpp b/src-hip/models.hip.cpp new file mode 100644 index 0000000000..b0848c971c --- /dev/null +++ b/src-hip/models.hip.cpp @@ -0,0 +1,855 @@ +#include "hip/hip_runtime.h" +// Removed vector dependency for MinGRU activations - now uses raw pointers + +#ifndef PUFFERLIB_MODELS_CU +#define PUFFERLIB_MODELS_CU + +#include +#include +#include + +#include "kernels.hip.cpp" + +// Signatures used by encoder and decoder. Writing custom nets in 4.0 requires a fair bit of code, +// because you are responsible for defining your own activation and gradient buffers. +// In practice, this is fairly simple. See our Encoder and Decoder for examples. +// You probably only ever need a custom Encoder +typedef void (*init_weights_fn)(void* weights, ulong* seed, hipStream_t stream); +typedef void (*reg_params_fn)(void* weights, Allocator* alloc); +typedef void (*reg_train_fn)(void* weights, void* buf, Allocator* acts, Allocator* grads, int B_TT); +typedef void (*reg_rollout_fn)(void* weights, void* buf, Allocator* alloc, int B); +typedef void* (*create_weights_fn)(void* self); +typedef void (*free_weights_fn)(void* weights); +typedef void (*free_activations_fn)(void* activations); +typedef PrecisionTensor (*forward_fn)(void* weights, void* activations, PrecisionTensor input, hipStream_t stream); +typedef void (*encoder_backward_fn)(void* weights, void* activations, + PrecisionTensor grad, hipStream_t stream); +typedef PrecisionTensor (*decoder_backward_fn)(void* weights, void* activations, + FloatTensor grad_logits, FloatTensor grad_logstd, FloatTensor grad_value, hipStream_t stream); +typedef PrecisionTensor (*network_forward_fn)(void* weights, PrecisionTensor x, + PrecisionTensor state, void* activations, hipStream_t stream); +typedef PrecisionTensor (*network_forward_train_fn)(void* weights, PrecisionTensor x, + PrecisionTensor state, void* activations, hipStream_t stream); +typedef PrecisionTensor (*network_backward_fn)(void* weights, + PrecisionTensor grad, void* activations, hipStream_t stream); + +struct Encoder { + forward_fn forward; + encoder_backward_fn backward; + init_weights_fn init_weights; + reg_params_fn reg_params; + reg_train_fn reg_train; + reg_rollout_fn reg_rollout; + create_weights_fn create_weights; + free_weights_fn free_weights; + free_activations_fn free_activations; + int in_dim, out_dim; + size_t activation_size; // sizeof(EncoderActivations) or custom override +}; + +struct Decoder { + forward_fn forward; + decoder_backward_fn backward; + init_weights_fn init_weights; + reg_params_fn reg_params; + reg_train_fn reg_train; + reg_rollout_fn reg_rollout; + create_weights_fn create_weights; + free_weights_fn free_weights; + free_activations_fn free_activations; + int hidden_dim, output_dim; + bool continuous; + int activation_size; // sizeof the impl's activations struct (custom decoders differ) +}; + +struct Network { + network_forward_fn forward; + network_forward_train_fn forward_train; + network_backward_fn backward; + init_weights_fn init_weights; + reg_params_fn reg_params; + reg_train_fn reg_train; + reg_rollout_fn reg_rollout; + create_weights_fn create_weights; + free_weights_fn free_weights; + free_activations_fn free_activations; + int hidden, num_layers, horizon; +}; + +struct EncoderWeights { + PrecisionTensor weight; + int in_dim, out_dim; +}; + +struct EncoderActivations { + PrecisionTensor out, saved_input, wgrad_scratch; +}; + +// The core of 4.0 is the MinGRU fused scan operation. This allows us to parallelize +// training across the sequence dimension and scale to longer sequences +__device__ __forceinline__ void log_coeffs_and_values_fwd(float gate, float hidden, + float* log_coeff_out, float* log_value_out) { + float abs_gate = fabsf(gate); + float sp_neg = log1pf(expf(-abs_gate)); + float softplus_gate = (gate >= 0.0f) ? gate + sp_neg : sp_neg; + float softplus_neg_gate = (gate >= 0.0f) ? sp_neg : -gate + sp_neg; + *log_coeff_out = -softplus_gate; + float log_tilde_h = (hidden >= 0.0f) ? logf(hidden + 0.5f) : -softplus_fwd(-hidden); + *log_value_out = -softplus_neg_gate + log_tilde_h; +} + +__device__ __forceinline__ void log_coeffs_and_values_bwd(float grad_log_coeffs, float grad_log_values, + float gate, float hidden, float* grad_gate_out, float* grad_hidden_out) { + float sig_gate = sigmoid(gate); + *grad_gate_out = -grad_log_coeffs * sig_gate + grad_log_values * (1.0f - sig_gate); + *grad_hidden_out = (hidden >= 0.0f) ? grad_log_values / (hidden + 0.5f) : grad_log_values * sigmoid(-hidden); +} + +__global__ void mingru_gate(precision_t* out, precision_t* next_state, + const precision_t* combined, const precision_t* state_in, + const precision_t* x_in, int H, int B) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int N = B * H; + if (idx >= N) { + return; + } + + int b = idx / H; + int h = idx % H; + + // combined = linear(x_in) = (B, H) -> (B, 3*H) + int combined_base = b * 3 * H; + float hidden = to_float(combined[combined_base + h]); + float gate = to_float(combined[combined_base + H + h]); + float proj = to_float(combined[combined_base + 2*H + h]); + float state = to_float(state_in[idx]); + float x = to_float(x_in[idx]); + + // mingru_gate computation + float gate_sigmoid = sigmoid(gate); + float hidden_tilde = (hidden >= 0.0f) ? hidden + 0.5f : fast_sigmoid(hidden); + float mingru_out = lerp(state, hidden_tilde, gate_sigmoid); + + // next_state is mingru_out (for recurrence) + next_state[idx] = from_float(mingru_out); + + // Highway connection: sigmoid(proj) * mingru_out + (1 - sigmoid(proj)) * x (highway gate) + float proj_sigmoid = sigmoid(proj); + out[idx] = from_float(proj_sigmoid * mingru_out + (1.0f - proj_sigmoid) * x); +} + +// Prefix scan buffers +struct PrefixScan { + precision_t* combined_ptr = nullptr; + precision_t* state_ptr = nullptr; + precision_t* input_ptr = nullptr; // (B, T, H) original input before projection (for highway gate) + int B = 0, T = 0, H = 0; + FloatTensor a_star, s_vals, log_values_buf; + PrecisionTensor out, next_state; + PrecisionTensor grad_combined, grad_state; + PrecisionTensor grad_input; // (B, T, H) highway gate gradient w.r.t. input +}; + +// Checkpointing trades off partial recomputation for memory bandwidth. +#define CHECKPOINT_INTERVAL 4 +__global__ void mingru_scan_forward(PrefixScan scan) { + int T_seq = scan.T, H = scan.H, B = scan.B; + precision_t* __restrict__ out = scan.out.data; + precision_t* __restrict__ next_state = scan.next_state.data; + float* __restrict__ a_star_buf = scan.a_star.data; + float* __restrict__ s_buf = scan.s_vals.data; + float* __restrict__ log_values_buf = scan.log_values_buf.data; + const precision_t* __restrict__ combined = scan.combined_ptr; + const precision_t* __restrict__ state = scan.state_ptr; + const precision_t* __restrict__ input = scan.input_ptr; + + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B * H) { + return; + } + + int b = idx / H; + int h = idx % H; + + int bH = b * H; + int H3 = 3 * H; + int H2 = 2 * H; + int bHT = bH * T_seq; + int out_base = bHT + h; + int cbase = 3 * bHT; + + float a_star = 0.0f; + float log_value = 0.0f; + + // Handle t=0 outside the loop: use log(state), coeff = 0 + float s = __logf(to_float(state[bH + h])); + log_value = s; + + int T_out = T_seq + 1; + int buf_base = b * T_out * H + h; + int buf_curr = buf_base; + a_star_buf[buf_curr] = a_star; + s_buf[buf_curr] = s; + log_values_buf[buf_curr] = log_value; + + const precision_t* combined_h_base = &combined[cbase + h]; + const precision_t* combined_g_base = &combined[cbase + H + h]; + const precision_t* combined_p_base = &combined[cbase + H2 + h]; + + // Loop t=1..T_seq with sparse checkpointing + float scan_result = 0.0f; + int out_curr = out_base; + int t_offset = 0; + + for (int t = 1; t < T_seq + 1; t++) { + float hidden_val = to_float(combined_h_base[t_offset]); + float gate_val = to_float(combined_g_base[t_offset]); + float proj_val = to_float(combined_p_base[t_offset]); + float x_val = to_float(input[out_base + (t - 1) * H]); + + float log_coeff_val; + log_coeffs_and_values_fwd(gate_val, hidden_val, &log_coeff_val, &log_value); + + // a_star[t] = sum_{i=0}^t log_coeffs[i] + a_star += log_coeff_val; + + float z = log_value - a_star; + s = logaddexp(s, z); + + scan_result = __expf(a_star + s); + float proj_sigmoid = sigmoid(proj_val); + + // out = sigmoid(proj) * scan_result + (1 - sigmoid(proj)) * x (highway gate) + out[out_curr] = from_float(proj_sigmoid * scan_result + (1.0f - proj_sigmoid) * x_val); + + buf_curr += H; + out_curr += H; + t_offset += H3; + + if (t % CHECKPOINT_INTERVAL == 0) { + a_star_buf[buf_curr] = a_star; + s_buf[buf_curr] = s; + log_values_buf[buf_curr] = log_value; + } + } + + // Write timestep T to next_state (raw scan_result, no proj, for recurrence) + next_state[bH + h] = from_float(scan_result); +} + +// Reads sparse checkpoints from forward pass, recomputes intermediate values in chunks +__global__ void mingru_scan_backward(PrefixScan scan, + const precision_t* __restrict__ grad_out, + const precision_t* __restrict__ grad_next_state) { + int T_seq = scan.T, H = scan.H, B = scan.B; + precision_t* __restrict__ grad_combined = scan.grad_combined.data; + precision_t* __restrict__ grad_state = scan.grad_state.data; + precision_t* __restrict__ grad_input = scan.grad_input.data; + const precision_t* __restrict__ combined = scan.combined_ptr; + const precision_t* __restrict__ state = scan.state_ptr; + const precision_t* __restrict__ input = scan.input_ptr; + const float* __restrict__ a_star_buf = scan.a_star.data; + const float* __restrict__ s_buf = scan.s_vals.data; + const float* __restrict__ log_values_buf = scan.log_values_buf.data; + + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B * H) { + return; + } + + int b = idx / H; + int h = idx % H; + + int bHT = b * H * T_seq; + int cbase = 3 * bHT; + int H3 = 3 * H; + int H2 = 2 * H; + const int state_idx = b * H + h; + const int out_base = bHT + h; + + const precision_t* combined_h_base = &combined[cbase + h]; + const precision_t* combined_g_base = &combined[cbase + H + h]; + const precision_t* combined_p_base = &combined[cbase + H2 + h]; + + precision_t* grad_combined_h_base = &grad_combined[cbase + h]; + precision_t* grad_combined_g_base = &grad_combined[cbase + H + h]; + precision_t* grad_combined_p_base = &grad_combined[cbase + H2 + h]; + + int T_out = T_seq + 1; + int buf_base = b * T_out * H + h; + + float acc = 0.0; + float s_val_next = 0.0; + float carry_grad_a = 0.0; + + for (int chunk_end = T_seq; chunk_end > 0; chunk_end -= CHECKPOINT_INTERVAL) { + int chunk_start = (chunk_end > CHECKPOINT_INTERVAL) ? (chunk_end - CHECKPOINT_INTERVAL) : 0; + int chunk_len = chunk_end - chunk_start; + + // Chunk storage in registers + float chunk_a_star[CHECKPOINT_INTERVAL]; + float chunk_s[CHECKPOINT_INTERVAL]; + float chunk_log_values[CHECKPOINT_INTERVAL]; + float chunk_hidden[CHECKPOINT_INTERVAL]; + float chunk_gate[CHECKPOINT_INTERVAL]; + + // Load checkpoint from global memory + int ckpt_buf_idx = buf_base + chunk_start * H; + float recomp_a_star = a_star_buf[ckpt_buf_idx]; + float recomp_s = s_buf[ckpt_buf_idx]; + float recomp_log_value = log_values_buf[ckpt_buf_idx]; + + // Recompute and store from chunk_start to chunk_end + for (int i = 0; i < chunk_len; ++i) { + int t = chunk_start + 1 + i; + int t_offset = (t - 1) * H3; + float hv = to_float(combined_h_base[t_offset]); + float gv = to_float(combined_g_base[t_offset]); + + float lc; + log_coeffs_and_values_fwd(gv, hv, &lc, &recomp_log_value); + recomp_a_star += lc; + + float z = recomp_log_value - recomp_a_star; + recomp_s = logaddexp(recomp_s, z); + + chunk_a_star[i] = recomp_a_star; + chunk_s[i] = recomp_s; + chunk_log_values[i] = recomp_log_value; + chunk_hidden[i] = hv; + chunk_gate[i] = gv; + } + + for (int i = chunk_len - 1; i >= 0; --i) { + int t = chunk_start + 1 + i; + int t_offset = (t - 1) * H3; + + float a_star_t = chunk_a_star[i]; + float s_t = chunk_s[i]; + float log_value_t = chunk_log_values[i]; + float hidden_val = chunk_hidden[i]; + float gate_val = chunk_gate[i]; + + float proj_val = to_float(combined_p_base[t_offset]); + int input_idx = out_base + (t - 1) * H; + float x_val = to_float(input[input_idx]); + + float scan_result = __expf(a_star_t + s_t); + float z = log_value_t - a_star_t; + + float grad_out_val = to_float(grad_out[input_idx]); + float grad_scan_from_next = (t == T_seq) ? to_float(grad_next_state[state_idx]) : 0.0f; + float proj_sigmoid = sigmoid(proj_val); + + // Highway gate gradients: out = sigmoid(proj) * scan_result + (1 - sigmoid(proj)) * x + float grad_scan_result = grad_scan_from_next + grad_out_val * proj_sigmoid; + float grad_proj = grad_out_val * (scan_result - x_val) * proj_sigmoid * (1.0f - proj_sigmoid); + grad_input[input_idx] = from_float(grad_out_val * (1.0f - proj_sigmoid)); + + float grad_log_h = grad_scan_result * scan_result; + float grad_s = grad_log_h; + + if (t == T_seq) { + acc = grad_s; + } else { + acc = grad_s + acc * __expf(s_t - s_val_next); + } + float grad_z = acc * __expf(z - s_t); + s_val_next = s_t; + + float grad_a = grad_log_h + carry_grad_a - grad_z; + carry_grad_a = grad_a; + + float grad_g, grad_h; + log_coeffs_and_values_bwd(grad_a, grad_z, gate_val, hidden_val, &grad_g, &grad_h); + + grad_combined_h_base[t_offset] = from_float(grad_h); + grad_combined_g_base[t_offset] = from_float(grad_g); + grad_combined_p_base[t_offset] = from_float(grad_proj); + } + } + + int ckpt_0_idx = buf_base; + float a_star_0 = a_star_buf[ckpt_0_idx]; + float s_0 = s_buf[ckpt_0_idx]; + float log_value_0 = log_values_buf[ckpt_0_idx]; + + float scan_result_0 = __expf(a_star_0 + s_0); + float z_0 = log_value_0 - a_star_0; + + float grad_scan_result_0 = 0.0f; + float grad_log_h_0 = grad_scan_result_0 * scan_result_0; + float grad_s_0 = grad_log_h_0; + + acc = grad_s_0 + acc * __expf(s_0 - s_val_next); + float grad_z_0 = acc * __expf(z_0 - s_0); + + grad_state[state_idx] = from_float(grad_z_0 / to_float(state[state_idx])); +} + +__global__ void sum_rows_to_precision_kernel(precision_t* __restrict__ dst, + const float* __restrict__ src, int R, int C) { + int col = blockIdx.x * blockDim.x + threadIdx.x; + if (col >= C) { + return; + } + float sum = 0.0f; + for (int r = 0; r < R; r++) { + sum += src[r * C + col]; + } + dst[col] = from_float(sum); +} + +__global__ void assemble_decoder_grad( + precision_t* __restrict__ dst, const float* __restrict__ grad_logits, + const float* __restrict__ grad_value, int B_TT, int od, int od_plus_1) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B_TT * od_plus_1) { + return; + } + int row = idx / od_plus_1, col = idx % od_plus_1; + dst[idx] = from_float((col < od) ? grad_logits[row * od + col] : grad_value[row]); +} + +static PrecisionTensor encoder_forward(void* w, void* activations, PrecisionTensor input, hipStream_t stream) { + EncoderWeights* ew = (EncoderWeights*)w; + EncoderActivations* a = (EncoderActivations*)activations; + if (a->saved_input.data) puf_copy(&a->saved_input, &input, stream); + puf_mm(&input, &ew->weight, &a->out, stream); + return a->out; +} + +static void encoder_backward(void* w, void* activations, PrecisionTensor grad, hipStream_t stream) { + EncoderActivations* a = (EncoderActivations*)activations; + puf_mm_tn(&grad, &a->saved_input, &a->wgrad_scratch, stream); +} + +static void encoder_init_weights(void* w, ulong* seed, hipStream_t stream) { + EncoderWeights* ew = (EncoderWeights*)w; + PrecisionTensor wt = { + .data = ew->weight.data, + .shape = {ew->out_dim, ew->in_dim}, + }; + puf_kaiming_init(&wt, std::sqrt(2.0f), (*seed)++, stream); +} + +static void encoder_reg_params(void* w, Allocator* alloc) { + EncoderWeights* ew = (EncoderWeights*)w; + ew->weight = {.shape = {ew->out_dim, ew->in_dim}}; + alloc_register(alloc,&ew->weight); +} + +static void encoder_reg_train(void* w, void* activations, Allocator* acts, Allocator* grads, int B_TT) { + EncoderWeights* ew = (EncoderWeights*)w; + EncoderActivations* a = (EncoderActivations*)activations; + *a = (EncoderActivations){ + .out = {.shape = {B_TT, ew->out_dim}}, + .saved_input = {.shape = {B_TT, ew->in_dim}}, + .wgrad_scratch = {.shape = {ew->out_dim, ew->in_dim}}, + }; + alloc_register(acts,&a->out); + alloc_register(acts,&a->saved_input); + alloc_register(grads,&a->wgrad_scratch); +} + +static void encoder_reg_rollout(void* w, void* activations, Allocator* alloc, int B) { + EncoderWeights* ew = (EncoderWeights*)w; + EncoderActivations* a = (EncoderActivations*)activations; + a->out = {.shape = {B, ew->out_dim}}; + alloc_register(alloc,&a->out); +} + +static void* encoder_create_weights(void* self) { + Encoder* e = (Encoder*)self; + EncoderWeights* ew = (EncoderWeights*)calloc(1, sizeof(EncoderWeights)); + ew->in_dim = e->in_dim; ew->out_dim = e->out_dim; + return ew; +} + +static void encoder_free_weights(void* weights) { + free(weights); +} + +static void encoder_free_activations(void* activations) { + free(activations); +} + +struct DecoderWeights { + PrecisionTensor weight, logstd; + int hidden_dim, output_dim; + bool continuous; +}; + +struct DecoderActivations { + PrecisionTensor out, grad_out, saved_input, grad_input, wgrad_scratch, logstd_scratch; +}; + +static PrecisionTensor decoder_forward(void* w, void* activations, PrecisionTensor input, hipStream_t stream) { + DecoderWeights* dw = (DecoderWeights*)w; + DecoderActivations* a = (DecoderActivations*)activations; + if (a->saved_input.data) { + puf_copy(&a->saved_input, &input, stream); + } + puf_mm(&input, &dw->weight, &a->out, stream); + return a->out; +} + +static void decoder_init_weights(void* w, ulong* seed, hipStream_t stream) { + DecoderWeights* dw = (DecoderWeights*)w; + PrecisionTensor wt = { + .data = dw->weight.data, + .shape = {dw->output_dim + 1, dw->hidden_dim}, + }; + puf_kaiming_init(&wt, 1.0f, (*seed)++, stream); +} + +static void decoder_reg_params(void* w, Allocator* alloc) { + DecoderWeights* dw = (DecoderWeights*)w; + dw->weight = {.shape = {dw->output_dim + 1, dw->hidden_dim}}; + alloc_register(alloc,&dw->weight); + if (dw->continuous) { + dw->logstd = {.shape = {1, dw->output_dim}}; + alloc_register(alloc,&dw->logstd); + } +} + +static void decoder_reg_train(void* w, void* activations, Allocator* acts, Allocator* grads, int B_TT) { + DecoderWeights* dw = (DecoderWeights*)w; + DecoderActivations* a = (DecoderActivations*)activations; + int od1 = dw->output_dim + 1; + *a = (DecoderActivations){ + .out = {.shape = {B_TT, od1}}, + .grad_out = {.shape = {B_TT, od1}}, + .saved_input = {.shape = {B_TT, dw->hidden_dim}}, + .grad_input = {.shape = {B_TT, dw->hidden_dim}}, + .wgrad_scratch = {.shape = {od1, dw->hidden_dim}}, + .logstd_scratch = {.shape = {1, dw->output_dim}}, + }; + alloc_register(acts,&a->out); + alloc_register(acts,&a->saved_input); + alloc_register(acts,&a->grad_out); + alloc_register(acts,&a->grad_input); + alloc_register(grads,&a->wgrad_scratch); + if (dw->continuous) alloc_register(grads,&a->logstd_scratch); +} + +static void decoder_reg_rollout(void* w, void* activations, Allocator* alloc, int B) { + DecoderWeights* dw = (DecoderWeights*)w; + DecoderActivations* a = (DecoderActivations*)activations; + a->out = {.shape = {B, dw->output_dim + 1}}; + alloc_register(alloc,&a->out); +} + +static void* decoder_create_weights(void* self) { + Decoder* d = (Decoder*)self; + DecoderWeights* dw = (DecoderWeights*)calloc(1, sizeof(DecoderWeights)); + dw->hidden_dim = d->hidden_dim; dw->output_dim = d->output_dim; dw->continuous = d->continuous; + return dw; +} + +static void decoder_free_weights(void* weights) { + free(weights); +} + +static void decoder_free_activations(void* activations) { + free(activations); +} + +static PrecisionTensor decoder_backward(void* w, void* activations, + FloatTensor grad_logits, FloatTensor grad_logstd, FloatTensor grad_value, hipStream_t stream) { + DecoderWeights* dw = (DecoderWeights*)w; + DecoderActivations* a = (DecoderActivations*)activations; + int B_TT = a->saved_input.shape[0]; + int od = dw->output_dim, od1 = od + 1; + assemble_decoder_grad<<>>( + a->grad_out.data, grad_logits.data, grad_value.data, B_TT, od, od1); + puf_mm_tn(&a->grad_out, &a->saved_input, &a->wgrad_scratch, stream); + if (dw->continuous && grad_logstd.data != nullptr) { + sum_rows_to_precision_kernel<<output_dim), BLOCK_SIZE, 0, stream>>>( + a->logstd_scratch.data, grad_logstd.data, B_TT, dw->output_dim); + } + puf_mm_nn(&a->grad_out, &dw->weight, &a->grad_input, stream); + return a->grad_input; +} + +struct MinGRUActivations { + int num_layers; + // Rollout + PrecisionTensor* combined; // (B rollout, 3*T)[num_layers] + PrecisionTensor out; // (B rollout, T) + PrecisionTensor next_state; // (B rollout, T) + // Training + PrecisionTensor* saved_inputs; // (B, TT, T)[num_layers] + PrefixScan* scan_bufs; // [num_layers] + PrecisionTensor* combined_bufs; // (B*TT, 3*T)[num_layers] + PrecisionTensor* wgrad_scratch; // (3*T, T)[num_layers] + PrecisionTensor grad_input_buf; // (B*TT, T) + PrecisionTensor grad_next_state; // (B, 1, T) +}; + +void mingru_activations_free(MinGRUActivations* a) { + free(a->combined); + free(a->saved_inputs); + free(a->scan_bufs); + free(a->combined_bufs); + free(a->wgrad_scratch); +} + +struct MinGRUWeights { + int hidden, num_layers, horizon; + PrecisionTensor* weights; // [num_layers] +}; + +static PrecisionTensor mingru_state_layer(MinGRUWeights* m, PrecisionTensor& state, int i) { + long B = state.shape[1], H = state.shape[2]; + return {.data = state.data + i * B * H, .shape = {B, H}}; +} + +static void mingru_init_weights(void* w, ulong* seed, hipStream_t stream) { + MinGRUWeights* m = (MinGRUWeights*)w; + for (int i = 0; i < m->num_layers; i++) { + PrecisionTensor w2d = { + .data = m->weights[i].data, + .shape = {3 * m->hidden, m->hidden}, + }; + puf_kaiming_init(&w2d, 1.0f, (*seed)++, stream); + } +} + +static void mingru_reg_params(void* w, Allocator* alloc) { + MinGRUWeights* m = (MinGRUWeights*)w; + for (int i = 0; i < m->num_layers; i++) { + m->weights[i] = {.shape = {3 * m->hidden, m->hidden}}; + alloc_register(alloc,&m->weights[i]); + } +} + +static void mingru_reg_train(void* w, void* activations, Allocator* acts, Allocator* grads, int B_TT) { + MinGRUWeights* m = (MinGRUWeights*)w; + MinGRUActivations* a = (MinGRUActivations*)activations; + int H = m->hidden, TT = m->horizon, B = B_TT / TT; + a->num_layers = m->num_layers; + a->saved_inputs = (PrecisionTensor*)calloc(m->num_layers, sizeof(PrecisionTensor)); + a->scan_bufs = (PrefixScan*)calloc(m->num_layers, sizeof(PrefixScan)); + a->combined_bufs = (PrecisionTensor*)calloc(m->num_layers, sizeof(PrecisionTensor)); + a->wgrad_scratch = (PrecisionTensor*)calloc(m->num_layers, sizeof(PrecisionTensor)); + a->grad_input_buf = {.shape = {B_TT, H}}; + a->grad_next_state = {.shape = {B, 1, H}}; + alloc_register(acts,&a->grad_input_buf); + alloc_register(acts,&a->grad_next_state); + for (int i = 0; i < m->num_layers; i++) { + a->scan_bufs[i] = { + .B = B, .T = TT, .H = H, + .a_star = {.shape = {B, TT + 1, H}}, + .s_vals = {.shape = {B, TT + 1, H}}, + .log_values_buf = {.shape = {B, TT + 1, H}}, + .out = {.shape = {B, TT, H}}, + .next_state = {.shape = {B, 1, H}}, + .grad_combined = {.shape = {B, TT, 3 * H}}, + .grad_state = {.shape = {B, 1, H}}, + .grad_input = {.shape = {B, TT, H}}, + }; + a->saved_inputs[i] = {.shape = {B, TT, H}}; + a->combined_bufs[i] = {.shape = {B_TT, 3 * H}}; + a->wgrad_scratch[i] = {.shape = {3 * H, H}}; + alloc_register(acts,&a->saved_inputs[i]); + alloc_register(acts,&a->combined_bufs[i]); + alloc_register(acts,&a->scan_bufs[i].out); + alloc_register(acts,&a->scan_bufs[i].next_state); + alloc_register(acts,&a->scan_bufs[i].a_star); + alloc_register(acts,&a->scan_bufs[i].s_vals); + alloc_register(acts,&a->scan_bufs[i].log_values_buf); + alloc_register(acts,&a->scan_bufs[i].grad_combined); + alloc_register(acts,&a->scan_bufs[i].grad_state); + alloc_register(acts,&a->scan_bufs[i].grad_input); + alloc_register(grads,&a->wgrad_scratch[i]); + } +} + +static void mingru_reg_rollout(void* weights, void* activations, Allocator* alloc, int B_inf) { + MinGRUWeights* w = (MinGRUWeights*)weights; + MinGRUActivations* a = (MinGRUActivations*)activations; + int H = w->hidden; + a->num_layers = w->num_layers; + a->combined = (PrecisionTensor*)calloc(w->num_layers, sizeof(PrecisionTensor)); + for (int i = 0; i < w->num_layers; i++) { + a->combined[i] = {.shape = {B_inf, 3 * H}}; + alloc_register(alloc,&a->combined[i]); + } + a->out = {.shape = {B_inf, H}}; + a->next_state = {.shape = {B_inf, H}}; + alloc_register(alloc,&a->out); + alloc_register(alloc,&a->next_state); +} + +static void* mingru_create_weights(void* self) { + Network* n = (Network*)self; + MinGRUWeights* mw = (MinGRUWeights*)calloc(1, sizeof(MinGRUWeights)); + mw->hidden = n->hidden; mw->num_layers = n->num_layers; mw->horizon = n->horizon; + mw->weights = (PrecisionTensor*)calloc(n->num_layers, sizeof(PrecisionTensor)); + return mw; +} + +static void mingru_free_weights(void* weights) { + MinGRUWeights* mw = (MinGRUWeights*)weights; + free(mw->weights); + free(mw); +} + +static void mingru_free_activations(void* activations) { + MinGRUActivations* a = (MinGRUActivations*)activations; + mingru_activations_free(a); + free(a); +} + +static PrecisionTensor mingru_forward(void* w, PrecisionTensor x, PrecisionTensor state, + void* activations, hipStream_t stream) { + MinGRUWeights* m = (MinGRUWeights*)w; + MinGRUActivations* a = (MinGRUActivations*)activations; + int B = state.shape[1]; + int H = state.shape[2]; + for (int i = 0; i < m->num_layers; i++) { + PrecisionTensor state_i = mingru_state_layer(m, state, i); + puf_mm(&x, &m->weights[i], &a->combined[i], stream); + mingru_gate<<>>( + a->out.data, a->next_state.data, + a->combined[i].data, state_i.data, x.data, H, B); + puf_copy(&state_i, &a->next_state, stream); + x = a->out; + } + return x; +} + +static PrecisionTensor mingru_forward_train(void* w, PrecisionTensor x, PrecisionTensor state, + void* activations, hipStream_t stream) { + MinGRUWeights* m = (MinGRUWeights*)w; + MinGRUActivations* a = (MinGRUActivations*)activations; + int B = x.shape[0]; + for (int i = 0; i < m->num_layers; i++) { + puf_copy(&a->saved_inputs[i], &x, stream); + PrecisionTensor state_i = mingru_state_layer(m, state, i); + puf_mm(&x, &m->weights[i], &a->combined_bufs[i], stream); + a->scan_bufs[i].combined_ptr = a->combined_bufs[i].data; + a->scan_bufs[i].state_ptr = state_i.data; + a->scan_bufs[i].input_ptr = a->saved_inputs[i].data; + mingru_scan_forward<<hidden), BLOCK_SIZE, 0, stream>>>(a->scan_bufs[i]); + x = a->scan_bufs[i].out; + } + return x; +} + +static PrecisionTensor mingru_backward(void* w, PrecisionTensor grad, void* activations, hipStream_t stream) { + MinGRUWeights* m = (MinGRUWeights*)w; + MinGRUActivations* a = (MinGRUActivations*)activations; + for (int i = m->num_layers - 1; i >= 0; i--) { + PrefixScan& scan = a->scan_bufs[i]; + mingru_scan_backward<<>>( + scan, grad.data, a->grad_next_state.data); + puf_mm_tn(&scan.grad_combined, &a->saved_inputs[i], &a->wgrad_scratch[i], stream); + puf_mm_nn(&scan.grad_combined, &m->weights[i], &a->grad_input_buf, stream); + int n = numel(scan.grad_input.shape); + add_kernel<<>>( + a->grad_input_buf.data, scan.grad_input.data, n); + grad = a->grad_input_buf; + } + return grad; +} + +struct Policy { + Encoder encoder; + Decoder decoder; + Network network; + int input_dim, hidden_dim, output_dim; + int num_atns; +}; + +struct PolicyActivations { + void* encoder; + void* decoder; + void* network; +}; + +struct PolicyWeights { + void* encoder; + void* decoder; + void* network; +}; + +static void policy_activations_free(Policy* p, PolicyActivations& a) { + p->encoder.free_activations(a.encoder); + p->decoder.free_activations(a.decoder); + p->network.free_activations(a.network); +} + +PrecisionTensor policy_forward(Policy* p, PolicyWeights& w, PolicyActivations& activations, + PrecisionTensor obs, PrecisionTensor state, hipStream_t stream) { + PrecisionTensor enc_out = p->encoder.forward(w.encoder, activations.encoder, obs, stream); + PrecisionTensor h = p->network.forward(w.network, enc_out, state, activations.network, stream); + return p->decoder.forward(w.decoder, activations.decoder, h, stream); +} + +PrecisionTensor policy_forward_train(Policy* p, PolicyWeights& w, PolicyActivations& activations, + PrecisionTensor x, PrecisionTensor state, hipStream_t stream) { + int B = x.shape[0], TT = x.shape[1]; + PrecisionTensor h = p->encoder.forward(w.encoder, activations.encoder, *puf_squeeze(&x, 0), stream); + h = p->network.forward_train(w.network, *puf_unsqueeze(&h, 0, B, TT), state, activations.network, stream); + PrecisionTensor dec_out = p->decoder.forward(w.decoder, activations.decoder, *puf_squeeze(&h, 0), stream); + return *puf_unsqueeze(&dec_out, 0, B, TT); +} + +void policy_backward(Policy* p, PolicyWeights& w, PolicyActivations& activations, + FloatTensor grad_logits, FloatTensor grad_logstd, FloatTensor grad_value, hipStream_t stream) { + int B = grad_logits.shape[0], TT = grad_logits.shape[1]; + PrecisionTensor grad_h = p->decoder.backward(w.decoder, activations.decoder, + *puf_squeeze(&grad_logits, 0), grad_logstd, *puf_squeeze(&grad_value, 0), stream); + grad_h = p->network.backward(w.network, *puf_unsqueeze(&grad_h, 0, B, TT), activations.network, stream); + p->encoder.backward(w.encoder, activations.encoder, grad_h, stream); +} + +PolicyActivations policy_reg_train(Policy* p, PolicyWeights& w, + Allocator* acts, Allocator* grads, int B_TT) { + PolicyActivations a; + a.encoder = calloc(1, p->encoder.activation_size); + a.decoder = calloc(1, p->decoder.activation_size); + a.network = calloc(1, sizeof(MinGRUActivations)); + p->encoder.reg_train(w.encoder, a.encoder, acts, grads, B_TT); + p->decoder.reg_train(w.decoder, a.decoder, acts, grads, B_TT); + p->network.reg_train(w.network, a.network, acts, grads, B_TT); + return a; +} + +PolicyActivations policy_reg_rollout(Policy* p, PolicyWeights& w, Allocator* acts, int B_inf) { + PolicyActivations a; + a.encoder = calloc(1, p->encoder.activation_size); + a.decoder = calloc(1, p->decoder.activation_size); + a.network = calloc(1, sizeof(MinGRUActivations)); + p->encoder.reg_rollout(w.encoder, a.encoder, acts, B_inf); + p->decoder.reg_rollout(w.decoder, a.decoder, acts, B_inf); + p->network.reg_rollout(w.network, a.network, acts, B_inf); + return a; +} + +void policy_init_weights(Policy* p, PolicyWeights& w, uint64_t* seed, hipStream_t stream) { + p->encoder.init_weights(w.encoder, seed, stream); + p->decoder.init_weights(w.decoder, seed, stream); + p->network.init_weights(w.network, seed, stream); +} + +PolicyWeights policy_weights_create(Policy* p, Allocator* params) { + PolicyWeights w; + w.encoder = p->encoder.create_weights(&p->encoder); + w.decoder = p->decoder.create_weights(&p->decoder); + w.network = p->network.create_weights(&p->network); + p->encoder.reg_params(w.encoder, params); + p->decoder.reg_params(w.decoder, params); + p->network.reg_params(w.network, params); + return w; +} + +void policy_weights_free(Policy* p, PolicyWeights* w) { + p->encoder.free_weights(w->encoder); + p->decoder.free_weights(w->decoder); + p->network.free_weights(w->network); +} + +#endif // PUFFERLIB_MODELS_CU diff --git a/src-hip/muon.hip.cpp b/src-hip/muon.hip.cpp new file mode 100644 index 0000000000..3fbed8775a --- /dev/null +++ b/src-hip/muon.hip.cpp @@ -0,0 +1,228 @@ +#include "hip/hip_runtime.h" +#include + +__global__ void muon_norm_reduce(float* __restrict__ out, const float* __restrict__ partials, int num_blocks) { + __shared__ float sdata[256]; + int tid = threadIdx.x; + sdata[tid] = (tid < num_blocks) ? partials[tid] : 0.0f; + __syncthreads(); + for (int s = blockDim.x / 2; s > 0; s >>= 1) { + if (tid < s) { + sdata[tid] += sdata[tid + s]; + } + __syncthreads(); + } + if (tid == 0) { + *out = sdata[0]; + } +} + +__global__ void muon_norm_partials(float* __restrict__ partials, const precision_t* __restrict__ src, int n) { + __shared__ float sdata[256]; + int tid = threadIdx.x; + float sum = 0.0f; + for (int i = blockIdx.x * blockDim.x + tid; i < n; i += blockDim.x * gridDim.x) { + float v = to_float(src[i]); + sum += v * v; + } + sdata[tid] = sum; + __syncthreads(); + for (int s = blockDim.x / 2; s > 0; s >>= 1) { + if (tid < s) { + sdata[tid] += sdata[tid + s]; + } + __syncthreads(); + } + if (tid == 0) { + partials[blockIdx.x] = sdata[0]; + } +} + +__global__ void muon_norm_apply(precision_t* __restrict__ dst, const float* __restrict__ norm_ptr, float eps, int n) { + float inv_norm = 1.0f / fmaxf(sqrtf(*norm_ptr), eps); + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = from_float(to_float(dst[idx]) * inv_norm); + } +} + +// Nesterov with f32 momentum accumulator and precision_t gradients +__global__ void muon_nesterov(float* __restrict__ mb, precision_t* __restrict__ gc, float mu, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + float m = mu * mb[idx] + to_float(gc[idx]); + mb[idx] = m; + gc[idx] = from_float(to_float(gc[idx]) + mu * m); + } +} + +// Fused weight update: wb = wb * (1 - lr*wd) - lr * scale * update +__global__ void muon_weight_update(float* __restrict__ wb, const precision_t* __restrict__ update, + const float* __restrict__ lr_ptr, float wd, float scale, int n) { + float lr = *lr_ptr; + float wd_scale = 1.0f - lr * wd; + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + wb[idx] = wb[idx] * wd_scale - lr * scale * to_float(update[idx]); + } +} + +__global__ void muon_clip_norm(precision_t* __restrict__ dst, + const float* __restrict__ sum_sq_ptr, float max_norm, float eps, int n) { + float clip_coef = fminf(max_norm / (sqrtf(*sum_sq_ptr) + eps), 1.0f); + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + dst[idx] = from_float(to_float(dst[idx]) * clip_coef); + } +} + +static constexpr double ns_coeffs[5][3] = { + {4.0848, -6.8946, 2.9270}, + {3.9505, -6.3029, 2.6377}, + {3.7418, -5.5913, 2.3037}, + {2.8769, -3.1427, 1.2046}, + {2.8366, -3.0525, 1.2012}, +}; + +struct Muon { + double momentum, weight_decay, eps; + float lr_val_init; + float* lr_ptr; + float* lr_derived_ptr; + float* norm_ptr; + float* grad_norm_ptr; + FloatTensor lr_puf, lr_derived_puf, ns_norm_puf, grad_norm_puf; + FloatTensor mb_puf; + PrecisionTensor gram, gram_buf, x_buf; + FloatTensor norm_partials; + long max_M, max_N; + Allocator* param_alloc; // params allocator — shapes used by muon_step + ncclComm_t nccl_comm; + int world_size; +}; + +void muon_init(Muon* m, Allocator* param_alloc, double lr_val, + double momentum, double eps, double weight_decay, + Allocator* alloc) { + m->momentum = momentum; + m->weight_decay = weight_decay; + m->eps = eps; + m->lr_val_init = (float)lr_val; + m->lr_ptr = nullptr; + m->lr_derived_ptr = nullptr; + m->param_alloc = param_alloc; + m->nccl_comm = nullptr; + m->world_size = 1; + m->max_M = 0; m->max_N = 0; + long n = param_alloc->total_elems; + m->lr_puf = {.shape = {1}}; + m->lr_derived_puf = {.shape = {2}}; + m->mb_puf = {.shape = {n}}; + m->norm_partials = {.shape = {256}}; + m->grad_norm_puf = {.shape = {1}}; + alloc_register(alloc, &m->lr_puf); + alloc_register(alloc, &m->lr_derived_puf); + alloc_register(alloc, &m->mb_puf); + alloc_register(alloc, &m->norm_partials); + alloc_register(alloc, &m->grad_norm_puf); + long max_M = 0, max_N = 0; + for (int _i = 0; _i < param_alloc->num_regs; _i++) { + AllocEntry& e = param_alloc->regs[_i]; + if (ndim(e.shape) >= 2) { + long R = e.shape[0], C = numel(e.shape) / R; + max_M = max(max_M, min(R, C)); + max_N = max(max_N, max(R, C)); + } + } + if (max_M > 0) { + m->max_M = max_M; m->max_N = max_N; + m->gram = {.shape = {max_M, max_M}}; + m->gram_buf = {.shape = {max_M, max_M}}; + m->x_buf = {.shape = {max_M, max_N}}; + m->ns_norm_puf = {.shape = {1}}; + alloc_register(alloc, &m->gram); + alloc_register(alloc, &m->gram_buf); + alloc_register(alloc, &m->x_buf); + alloc_register(alloc, &m->ns_norm_puf); + } +} + +void muon_post_create(Muon* m) { + m->lr_ptr = m->lr_puf.data; + m->lr_derived_ptr = m->lr_derived_puf.data; + m->grad_norm_ptr = m->grad_norm_puf.data; + if (m->ns_norm_puf.data) m->norm_ptr = m->ns_norm_puf.data; + hipMemcpy(m->lr_ptr, &m->lr_val_init, sizeof(float), hipMemcpyHostToDevice); + hipMemset(m->lr_derived_ptr, 0, 2 * sizeof(float)); + hipMemset(m->mb_puf.data, 0, numel(m->mb_puf.shape) * sizeof(float)); +} + +void muon_step(Muon* m, FloatTensor weights, PrecisionTensor grads, float max_grad_norm, hipStream_t stream = 0) { + // Multi-GPU support: simple all-reduce over a contiguous grad buffer + if (m->nccl_comm != nullptr && m->world_size > 1) { + ncclAllReduce(grads.data, grads.data, numel(grads.shape), + NCCL_PRECISION, ncclAvg, m->nccl_comm, stream); + } + + // Clip gradients by norm + int clip_blocks = min((int)grid_size(numel(grads.shape)), 256); + muon_norm_partials<<>>( + m->norm_partials.data, grads.data, numel(grads.shape)); + muon_norm_reduce<<<1, 256, 0, stream>>>(m->grad_norm_ptr, m->norm_partials.data, clip_blocks); + muon_clip_norm<<>>( + grads.data, m->grad_norm_ptr, max_grad_norm, 1e-6f, numel(grads.shape)); + + // Nesterov momentum + muon_nesterov<<mb_puf.shape)), BLOCK_SIZE, 0, stream>>>( + m->mb_puf.data, grads.data, (float)m->momentum, numel(m->mb_puf.shape)); + + long offset = 0; + for (int _i = 0; _i < m->param_alloc->num_regs; _i++) { + AllocEntry& e = m->param_alloc->regs[_i]; + precision_t* gc_ptr = grads.data + offset; + float* wb_ptr = weights.data + offset; + long ne = numel(e.shape); + const precision_t* update_ptr = gc_ptr; + float scale = 1.0f; + + // Orthogonalize the update + if (ndim(e.shape) >= 2) { + long R = e.shape[0], C = ne / R; + long M = min(R, C), N = max(R, C); + bool tall = R > C; + PrecisionTensor x = {.data = gc_ptr, .shape = {R, C}}; + PrecisionTensor x_buf = {.data = m->x_buf.data, .shape = {R, C}}; + PrecisionTensor gram = {.data = m->gram.data, .shape = {M, M}}; + PrecisionTensor gram_buf = {.data = m->gram_buf.data, .shape = {M, M}}; + + int nblk = min((int)grid_size(numel(x.shape)), 256); + muon_norm_partials<<>>( + m->norm_partials.data, x.data, numel(x.shape)); + muon_norm_reduce<<<1, 256, 0, stream>>>(m->norm_ptr, m->norm_partials.data, nblk); + muon_norm_apply<<>>( + x.data, m->norm_ptr, 1e-7f, numel(x.shape)); + + hipblasOperation_t gram_op_a = tall ? HIPBLAS_OP_T : HIPBLAS_OP_N; + hipblasOperation_t gram_op_b = tall ? HIPBLAS_OP_N : HIPBLAS_OP_T; + for (int i = 0; i < 5; ++i) { + PrecisionTensor& src = (i % 2 == 0) ? x : x_buf; + PrecisionTensor& dst = (i % 2 == 0) ? x_buf : x; + cublasGemmExDense(gram_op_a, gram_op_b, (int)M, (int)M, (int)N, + src.data, src.data, gram.data, stream); + puf_copy(&gram_buf, &gram, stream); + puf_addmm_nn(&gram, &gram, &gram_buf, ns_coeffs[i][2], ns_coeffs[i][1], stream); + puf_copy(&dst, &src, stream); + cublasGemmExDense(HIPBLAS_OP_N, HIPBLAS_OP_N, (int)R, (int)C, (int)M, + tall ? src.data : gram_buf.data, tall ? gram_buf.data : src.data, dst.data, + stream, 1.0f, ns_coeffs[i][0]); + } + + update_ptr = x_buf.data; + scale = sqrtf(fmaxf(1.0f, (float)R / (float)C)); + } + + muon_weight_update<<>>( + wb_ptr, update_ptr, m->lr_ptr, (float)m->weight_decay, scale, (int)ne); + offset += ne; + } +} diff --git a/src-hip/ocean.hip.cpp b/src-hip/ocean.hip.cpp new file mode 100644 index 0000000000..eed552217e --- /dev/null +++ b/src-hip/ocean.hip.cpp @@ -0,0 +1,599 @@ +#include "hip/hip_runtime.h" +// NMMO3 CUDA encoder: multihot, cuDNN conv, embedding, concat, projection +// Included by pufferlib.cu — requires precision_t, PrecisionTensor, Allocator, puf_mm, etc. + +#include "cudnn_conv2d.hip.cpp" +// nethack.cu is not part of the HIP port + +// ---- NMMO3 constants ---- + +static constexpr int N3_MAP_H = 11, N3_MAP_W = 15, N3_NFEAT = 10; +static constexpr int N3_MULTIHOT = 59; +static constexpr int N3_MAP_SIZE = N3_MAP_H * N3_MAP_W * N3_NFEAT; +static constexpr int N3_PLAYER = 47, N3_REWARD = 10; +static constexpr int N3_EMBED_DIM = 32, N3_EMBED_VOCAB = 128; +static constexpr int N3_PLAYER_EMBED = N3_PLAYER * N3_EMBED_DIM; +static constexpr int N3_C1_IC = 59, N3_C1_OC = 128, N3_C1_K = 5, N3_C1_S = 3; +static constexpr int N3_C1_OH = 3, N3_C1_OW = 4; +static constexpr int N3_C2_IC = 128, N3_C2_OC = 128, N3_C2_K = 3, N3_C2_S = 1; +static constexpr int N3_C2_OH = 1, N3_C2_OW = 2; +static constexpr int N3_CONV_FLAT = N3_C2_OC * N3_C2_OH * N3_C2_OW; +static constexpr int N3_CONCAT = N3_CONV_FLAT + N3_PLAYER_EMBED + N3_PLAYER + N3_REWARD; + +__constant__ int N3_OFFSETS[10] = {0, 4, 8, 25, 30, 33, 38, 43, 48, 55}; + +static miopenDataType_t n3_cudnn_dtype() { + assert(PRECISION_SIZE == 4 && "MIOpen conv path supports float precision only"); + return miopenFloat; +} + +// ---- NMMO3 kernels ---- + +__global__ void n3_multihot_kernel( + precision_t* __restrict__ out, const precision_t* __restrict__ obs, int B, int obs_size) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B * N3_MAP_H * N3_MAP_W) return; + int b = idx / (N3_MAP_H * N3_MAP_W), rem = idx % (N3_MAP_H * N3_MAP_W); + int h = rem / N3_MAP_W, w = rem % N3_MAP_W; + const precision_t* src = obs + b * obs_size + (h * N3_MAP_W + w) * N3_NFEAT; + precision_t* dst = out + b * N3_MULTIHOT * N3_MAP_H * N3_MAP_W; + for (int f = 0; f < N3_NFEAT; f++) + dst[(N3_OFFSETS[f] + (int)to_float(src[f])) * N3_MAP_H * N3_MAP_W + h * N3_MAP_W + w] = from_float(1.0f); +} + +__global__ void n3_embedding_kernel( + precision_t* __restrict__ out, const precision_t* __restrict__ obs, + const precision_t* __restrict__ embed_w, int B, int obs_size) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B * N3_PLAYER) return; + int b = idx / N3_PLAYER, f = idx % N3_PLAYER; + int val = (int)to_float(obs[b * obs_size + N3_MAP_SIZE + f]); + const precision_t* src = embed_w + val * N3_EMBED_DIM; + precision_t* dst = out + b * N3_PLAYER_EMBED + f * N3_EMBED_DIM; + for (int d = 0; d < N3_EMBED_DIM; d++) dst[d] = src[d]; +} + +__global__ void n3_concat_kernel( + precision_t* __restrict__ out, const precision_t* __restrict__ conv_flat, + const precision_t* __restrict__ embed, const precision_t* __restrict__ obs, + int B, int obs_size) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B * N3_CONCAT) return; + int b = idx / N3_CONCAT, c = idx % N3_CONCAT; + precision_t val; + if (c < N3_CONV_FLAT) { + int oc = c / (N3_C2_OH * N3_C2_OW), r = c % (N3_C2_OH * N3_C2_OW); + int oh = r / N3_C2_OW, ow = r % N3_C2_OW; + val = conv_flat[b * N3_CONV_FLAT + oc * N3_C2_OH * N3_C2_OW + oh * N3_C2_OW + ow]; + } else if (c < N3_CONV_FLAT + N3_PLAYER_EMBED) + val = embed[b * N3_PLAYER_EMBED + (c - N3_CONV_FLAT)]; + else if (c < N3_CONV_FLAT + N3_PLAYER_EMBED + N3_PLAYER) + val = obs[b * obs_size + N3_MAP_SIZE + (c - N3_CONV_FLAT - N3_PLAYER_EMBED)]; + else + val = obs[b * obs_size + obs_size - N3_REWARD + (c - N3_CONV_FLAT - N3_PLAYER_EMBED - N3_PLAYER)]; + out[idx] = val; +} + +__global__ void n3_bias_relu_kernel( + precision_t* __restrict__ data, const precision_t* __restrict__ bias, int total, int dim) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= total) return; + data[idx] = from_float(fmaxf(0.0f, to_float(data[idx]) + to_float(bias[idx % dim]))); +} + +__global__ void n3_relu_backward_kernel( + precision_t* __restrict__ grad, const precision_t* __restrict__ out, int total) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= total) return; + if (to_float(out[idx]) <= 0.0f) grad[idx] = from_float(0.0f); +} + + +__global__ void bias_grad_kernel( + precision_t* __restrict__ bgrad, const precision_t* __restrict__ grad, int N, int dim) { + int d = blockIdx.x; + if (d >= dim) return; + float sum = 0.0f; + for (int i = threadIdx.x; i < N; i += blockDim.x) + sum += to_float(grad[i * dim + d]); + for (int offset = 16; offset > 0; offset >>= 1) + sum += __shfl_down_sync(0xffffffffull, sum, offset); + __shared__ float sdata[32]; + int lane = threadIdx.x % 32, warp = threadIdx.x / 32; + if (lane == 0) sdata[warp] = sum; + __syncthreads(); + if (warp == 0) { + sum = (lane < (blockDim.x + 31) / 32) ? sdata[lane] : 0.0f; + for (int offset = 16; offset > 0; offset >>= 1) + sum += __shfl_down_sync(0xffffffffull, sum, offset); + if (lane == 0) bgrad[d] = from_float(sum); + } +} + +// NCHW bias grad: sum over (B, OH, OW) for each OC channel +__global__ void n3_conv_bias_grad_nchw( + precision_t* __restrict__ bgrad, const precision_t* __restrict__ grad, + int B, int OC, int spatial) { + int oc = blockIdx.x; + if (oc >= OC) return; + float sum = 0.0f; + int total = B * spatial; + for (int i = threadIdx.x; i < total; i += blockDim.x) { + int b = i / spatial, s = i % spatial; + sum += to_float(grad[b * OC * spatial + oc * spatial + s]); + } + for (int offset = 16; offset > 0; offset >>= 1) + sum += __shfl_down_sync(0xffffffffull, sum, offset); + __shared__ float sdata[32]; + int lane = threadIdx.x % 32, warp = threadIdx.x / 32; + if (lane == 0) sdata[warp] = sum; + __syncthreads(); + if (warp == 0) { + sum = (lane < (blockDim.x + 31) / 32) ? sdata[lane] : 0.0f; + for (int offset = 16; offset > 0; offset >>= 1) + sum += __shfl_down_sync(0xffffffffull, sum, offset); + if (lane == 0) bgrad[oc] = from_float(sum); + } +} + +__global__ void n3_concat_backward_conv_kernel( + precision_t* __restrict__ conv_grad, const precision_t* __restrict__ concat_grad, int B) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B * N3_CONV_FLAT) return; + int b = idx / N3_CONV_FLAT, c = idx % N3_CONV_FLAT; + conv_grad[b * N3_CONV_FLAT + c] = concat_grad[b * N3_CONCAT + c]; +} + +// Embedding backward: scatter-add grad from concat_grad's player_embed region +// into embed_wgrad (float accumulation buffer). +// Each (b, f) looked up row obs[b, MAP_SIZE+f] from the table. +__global__ void n3_embedding_backward_kernel( + float* __restrict__ embed_wgrad_f, + const precision_t* __restrict__ concat_grad, + const precision_t* __restrict__ obs, + int B, int obs_size) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B * N3_PLAYER * N3_EMBED_DIM) return; + int b = idx / (N3_PLAYER * N3_EMBED_DIM); + int rem = idx % (N3_PLAYER * N3_EMBED_DIM); + int f = rem / N3_EMBED_DIM; + int d = rem % N3_EMBED_DIM; + int val = (int)to_float(obs[b * obs_size + N3_MAP_SIZE + f]); + float g = to_float(concat_grad[b * N3_CONCAT + N3_CONV_FLAT + f * N3_EMBED_DIM + d]); + atomicAdd(&embed_wgrad_f[val * N3_EMBED_DIM + d], g); +} + +// Cast float buffer to precision_t +__global__ void n3_float_to_precision_kernel( + precision_t* __restrict__ dst, const float* __restrict__ src, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) dst[idx] = from_float(src[idx]); +} + +// ---- atomicAdd for precision_t ---- +#ifdef PRECISION_FLOAT +__device__ __forceinline__ void atomicAdd_precision(precision_t* addr, precision_t val) { + atomicAdd(addr, val); +} +#else +__device__ __forceinline__ void atomicAdd_precision(precision_t* addr, precision_t val) { + // bf16 atomicAdd via CAS on enclosing 32-bit word + unsigned int* addr_u32 = (unsigned int*)((size_t)addr & ~2ULL); + bool is_high = ((size_t)addr & 2) != 0; + unsigned int old_u32 = *addr_u32, assumed; + do { + assumed = old_u32; + __hip_bfloat16* pair = (__hip_bfloat16*)&old_u32; + float sum = __bfloat162float(pair[is_high]) + __bfloat162float(val); + unsigned int new_u32 = assumed; + ((__hip_bfloat16*)&new_u32)[is_high] = __float2bfloat16(sum); + old_u32 = atomicCAS(addr_u32, assumed, new_u32); + } while (old_u32 != assumed); +} +#endif + +// ---- NCHW bias kernels for im2col conv path ---- + +__global__ void conv_bias_kernel(precision_t* __restrict__ data, + const precision_t* __restrict__ bias, int B, int OC, int spatial) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = B * OC * spatial; + if (idx >= total) return; + int oc = (idx / spatial) % OC; + data[idx] = from_float(to_float(data[idx]) + to_float(bias[oc])); +} + +__global__ void conv_bias_relu_kernel(precision_t* __restrict__ data, + const precision_t* __restrict__ bias, int B, int OC, int spatial) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = B * OC * spatial; + if (idx >= total) return; + int oc = (idx / spatial) % OC; + data[idx] = from_float(fmaxf(0.0f, to_float(data[idx]) + to_float(bias[oc]))); +} + +// ---- im2col + cuBLAS conv (no cuDNN) ---- +// NCHW layout throughout. Weight stored as (OC, IC*K*K). +// im2col produces (B*OH*OW, IC*K*K), matmul with W^T gives (B*OH*OW, OC), +// then reshape to NCHW (B, OC, OH, OW). + +__global__ void im2col_kernel( + const precision_t* __restrict__ input, precision_t* __restrict__ col, + int B, int IC, int IH, int IW, int K, int S, int OH, int OW +) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = B * OH * OW * IC * K * K; + if (idx >= total) return; + int col_w = IC * K * K; + int row = idx / col_w; + int c = idx % col_w; + int b = row / (OH * OW); + int rem = row % (OH * OW); + int oh = rem / OW, ow = rem % OW; + int ic = c / (K * K), kk = c % (K * K); + int kh = kk / K, kw = kk % K; + int ih = oh * S + kh, iw = ow * S + kw; + col[idx] = input[b * IC * IH * IW + ic * IH * IW + ih * IW + iw]; +} + +// Backward: col2im — input-centric gather to avoid atomics. +// Each thread owns one (b, ic, ih, iw) element and sums contributions from all +// (oh, ow, kh, kw) patches that map to it. +__global__ void col2im_kernel( + const precision_t* __restrict__ col, precision_t* __restrict__ grad_input, + int B, int IC, int IH, int IW, int K, int S, int OH, int OW +) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = B * IC * IH * IW; + if (idx >= total) return; + int iw = idx % IW; + int ih = (idx / IW) % IH; + int ic = (idx / (IW * IH)) % IC; + int b = idx / (IW * IH * IC); + float sum = 0.0f; + for (int kh = 0; kh < K; kh++) { + int ih_off = ih - kh; + if (ih_off < 0 || ih_off % S != 0) continue; + int oh = ih_off / S; + if (oh >= OH) continue; + for (int kw = 0; kw < K; kw++) { + int iw_off = iw - kw; + if (iw_off < 0 || iw_off % S != 0) continue; + int ow = iw_off / S; + if (ow >= OW) continue; + int col_idx = (b * OH * OW + oh * OW + ow) * (IC * K * K) + ic * K * K + kh * K + kw; + sum += to_float(col[col_idx]); + } + } + grad_input[idx] = from_float(sum); +} + +// Transpose (B, OC, OH, OW) -> (B*OH*OW, OC) [NCHW to row-major spatial-first] +__global__ void nchw_to_rows_kernel( + const precision_t* __restrict__ src, precision_t* __restrict__ dst, + int B, int OC, int spatial +) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = B * OC * spatial; + if (idx >= total) return; + int b = idx / (OC * spatial); + int oc = (idx / spatial) % OC; + int s = idx % spatial; + dst[(b * spatial + s) * OC + oc] = src[idx]; +} + +// Transpose (B*OH*OW, OC) -> (B, OC, OH, OW) [row-major spatial-first to NCHW] +__global__ void rows_to_nchw_kernel( + const precision_t* __restrict__ src, precision_t* __restrict__ dst, + int B, int OC, int spatial +) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = B * OC * spatial; + if (idx >= total) return; + int b = idx / (OC * spatial); + int oc = (idx / spatial) % OC; + int s = idx % spatial; + dst[idx] = src[(b * spatial + s) * OC + oc]; +} + +// Forward: im2col conv + bias + optional relu. All NCHW. +// col_buf: pre-allocated (max_B * OH * OW, IC * K * K) +// mm_buf: pre-allocated (max_B * OH * OW, OC) — row-major (spatial-first) +static void gemm_conv_forward( + PrecisionTensor* weight, PrecisionTensor* bias, + precision_t* input, precision_t* output, + precision_t* col_buf, precision_t* mm_buf, + int B, int IC, int IH, int IW, int OC, int K, int S, int OH, int OW, + bool relu, hipStream_t stream +) { + int col_rows = B * OH * OW; + int col_cols = IC * K * K; + int total_col = col_rows * col_cols; + int total_out = B * OC * OH * OW; + + // im2col: input NCHW -> col (B*OH*OW, IC*K*K) + im2col_kernel<<>>( + input, col_buf, B, IC, IH, IW, K, S, OH, OW); + + // matmul: col (B*OH*OW, IC*K*K) @ W^T (IC*K*K, OC) = mm_buf (B*OH*OW, OC) + PrecisionTensor col_t = {.data = col_buf, .shape = {col_rows, col_cols}}; + PrecisionTensor mm_t = {.data = mm_buf, .shape = {col_rows, OC}}; + puf_mm(&col_t, weight, &mm_t, stream); + + // transpose (B*OH*OW, OC) -> (B, OC, OH, OW) NCHW + bias + relu + int spatial = OH * OW; + rows_to_nchw_kernel<<>>( + mm_buf, output, B, OC, spatial); + if (relu) { + conv_bias_relu_kernel<<>>( + output, bias->data, B, OC, spatial); + } else { + conv_bias_kernel<<>>( + output, bias->data, B, OC, spatial); + } +} + +// Backward: weight grad + optional input grad via im2col/col2im + cuBLAS. +// grad_output is NCHW (B, OC, OH, OW). saved_input is NCHW. +// Caller handles relu backward and bias grad (same as cuDNN path). +static void gemm_conv_backward( + PrecisionTensor* weight, + precision_t* saved_input, precision_t* grad_output, + precision_t* wgrad, precision_t* input_grad, + precision_t* col_buf, precision_t* mm_buf, + int B, int IC, int IH, int IW, int OC, int K, int S, int OH, int OW, + hipStream_t stream +) { + int col_rows = B * OH * OW; + int col_cols = IC * K * K; + int total_col = col_rows * col_cols; + int total_out = B * OC * OH * OW; + int spatial = OH * OW; + + // Transpose grad_output NCHW -> (B*OH*OW, OC) + nchw_to_rows_kernel<<>>( + grad_output, mm_buf, B, OC, spatial); + + // im2col of saved_input + im2col_kernel<<>>( + saved_input, col_buf, B, IC, IH, IW, K, S, OH, OW); + + // Weight grad: mm_buf^T (OC, B*OH*OW) @ col_buf (B*OH*OW, IC*K*K) = wgrad (OC, IC*K*K) + PrecisionTensor mm_t = {.data = mm_buf, .shape = {col_rows, OC}}; + PrecisionTensor col_t = {.data = col_buf, .shape = {col_rows, col_cols}}; + PrecisionTensor wg_t = {.data = wgrad, .shape = {OC, col_cols}}; + puf_mm_tn(&mm_t, &col_t, &wg_t, stream); + + // Input grad (optional): mm_buf (B*OH*OW, OC) @ weight (OC, IC*K*K) = col_grad (B*OH*OW, IC*K*K) + if (input_grad) { + puf_mm_nn(&mm_t, weight, &col_t, stream); // reuse col_buf as col_grad + col2im_kernel<<>>( + col_buf, input_grad, B, IC, IH, IW, K, S, OH, OW); + } +} + +// ---- NMMO3 encoder structs ---- + +struct NMMO3EncoderWeights { + ConvWeights conv1, conv2; + PrecisionTensor embed_w, proj_w, proj_b; + int obs_size, hidden; +}; + +struct NMMO3EncoderActivations { + ConvActivations conv1, conv2; + PrecisionTensor col1, mm1, col2, mm2; // im2col + matmul scratch buffers + PrecisionTensor multihot, embed_out, concat, out, saved_obs; + PrecisionTensor embed_wgrad, proj_wgrad, proj_bgrad; + FloatTensor embed_wgrad_f; // float accumulation buffer for scatter-add +}; + +static NMMO3EncoderWeights* nmmo3_encoder_create(int obs_size, int hidden) { + NMMO3EncoderWeights* ew = (NMMO3EncoderWeights*)calloc(1, sizeof(NMMO3EncoderWeights)); + ew->obs_size = obs_size; ew->hidden = hidden; + conv_init(&ew->conv1, N3_C1_IC, N3_C1_OC, N3_C1_K, N3_C1_S, N3_MAP_H, N3_MAP_W, true); + conv_init(&ew->conv2, N3_C2_IC, N3_C2_OC, N3_C2_K, N3_C2_S, N3_C1_OH, N3_C1_OW, false); + return ew; +} + +// ---- NMMO3 encoder interface ---- + +static PrecisionTensor nmmo3_encoder_forward(void* w, void* activations, PrecisionTensor input, hipStream_t stream) { + NMMO3EncoderWeights* ew = (NMMO3EncoderWeights*)w; + NMMO3EncoderActivations* a = (NMMO3EncoderActivations*)activations; + int B = input.shape[0]; + + if (a->saved_obs.data) puf_copy(&a->saved_obs, &input, stream); + + hipMemsetAsync(a->multihot.data, 0, (int64_t)B * N3_MULTIHOT * N3_MAP_H * N3_MAP_W * sizeof(precision_t), stream); + n3_multihot_kernel<<>>( + a->multihot.data, input.data, B, ew->obs_size); + + gemm_conv_forward(&ew->conv1.w, &ew->conv1.b, a->multihot.data, a->conv1.out.data, + a->col1.data, a->mm1.data, B, N3_C1_IC, N3_MAP_H, N3_MAP_W, + N3_C1_OC, N3_C1_K, N3_C1_S, N3_C1_OH, N3_C1_OW, true, stream); + if (a->conv1.saved_input.data) + hipMemcpyAsync(a->conv1.saved_input.data, a->multihot.data, + (int64_t)B * N3_C1_IC * N3_MAP_H * N3_MAP_W * sizeof(precision_t), hipMemcpyDeviceToDevice, stream); + gemm_conv_forward(&ew->conv2.w, &ew->conv2.b, a->conv1.out.data, a->conv2.out.data, + a->col2.data, a->mm2.data, B, N3_C2_IC, N3_C1_OH, N3_C1_OW, + N3_C2_OC, N3_C2_K, N3_C2_S, N3_C2_OH, N3_C2_OW, false, stream); + if (a->conv2.saved_input.data) + hipMemcpyAsync(a->conv2.saved_input.data, a->conv1.out.data, + (int64_t)B * N3_C2_IC * N3_C1_OH * N3_C1_OW * sizeof(precision_t), hipMemcpyDeviceToDevice, stream); + + n3_embedding_kernel<<>>( + a->embed_out.data, input.data, ew->embed_w.data, B, ew->obs_size); + n3_concat_kernel<<>>( + a->concat.data, a->conv2.out.data, a->embed_out.data, input.data, B, ew->obs_size); + + puf_mm(&a->concat, &ew->proj_w, &a->out, stream); + n3_bias_relu_kernel<<hidden), BLOCK_SIZE, 0, stream>>>( + a->out.data, ew->proj_b.data, B * ew->hidden, ew->hidden); + return a->out; +} + +static void nmmo3_encoder_backward(void* w, void* activations, PrecisionTensor grad, hipStream_t stream) { + NMMO3EncoderWeights* ew = (NMMO3EncoderWeights*)w; + NMMO3EncoderActivations* a = (NMMO3EncoderActivations*)activations; + int B = grad.shape[0], H = ew->hidden; + + n3_relu_backward_kernel<<>>( + grad.data, a->out.data, B * H); + bias_grad_kernel<<>>( + a->proj_bgrad.data, grad.data, B, H); + puf_mm_tn(&grad, &a->concat, &a->proj_wgrad, stream); + + PrecisionTensor grad_concat = {.data = a->concat.data, .shape = {B, N3_CONCAT}}; + puf_mm_nn(&grad, &ew->proj_w, &grad_concat, stream); + + n3_concat_backward_conv_kernel<<>>( + a->conv2.grad.data, grad_concat.data, B); + + n3_conv_bias_grad_nchw<<conv2.OC, 256, 0, stream>>>( + a->conv2.bgrad.data, a->conv2.grad.data, + B, ew->conv2.OC, ew->conv2.OH * ew->conv2.OW); + gemm_conv_backward(&ew->conv2.w, a->conv2.saved_input.data, a->conv2.grad.data, + a->conv2.wgrad.data, a->conv1.grad.data, + a->col2.data, a->mm2.data, B, N3_C2_IC, N3_C1_OH, N3_C1_OW, + N3_C2_OC, N3_C2_K, N3_C2_S, N3_C2_OH, N3_C2_OW, stream); + + n3_relu_backward_kernel<<conv1.OC * ew->conv1.OH * ew->conv1.OW), BLOCK_SIZE, 0, stream>>>( + a->conv1.grad.data, a->conv1.out.data, + B * ew->conv1.OC * ew->conv1.OH * ew->conv1.OW); + n3_conv_bias_grad_nchw<<conv1.OC, 256, 0, stream>>>( + a->conv1.bgrad.data, a->conv1.grad.data, + B, ew->conv1.OC, ew->conv1.OH * ew->conv1.OW); + gemm_conv_backward(&ew->conv1.w, a->conv1.saved_input.data, a->conv1.grad.data, + a->conv1.wgrad.data, NULL, + a->col1.data, a->mm1.data, B, N3_C1_IC, N3_MAP_H, N3_MAP_W, + N3_C1_OC, N3_C1_K, N3_C1_S, N3_C1_OH, N3_C1_OW, stream); + + // Embedding backward: scatter-add from concat gradient into float buffer, then cast + int embed_n = N3_EMBED_VOCAB * N3_EMBED_DIM; + hipMemsetAsync(a->embed_wgrad_f.data, 0, embed_n * sizeof(float), stream); + n3_embedding_backward_kernel<<>>( + a->embed_wgrad_f.data, grad_concat.data, a->saved_obs.data, B, ew->obs_size); + n3_float_to_precision_kernel<<>>( + a->embed_wgrad.data, a->embed_wgrad_f.data, embed_n); +} + +static void nmmo3_encoder_init_weights(void* w, uint64_t* seed, hipStream_t stream) { + NMMO3EncoderWeights* ew = (NMMO3EncoderWeights*)w; + conv_init_weights(&ew->conv1, seed, stream); + conv_init_weights(&ew->conv2, seed, stream); + auto init2d = [&](PrecisionTensor& t, int rows, int cols, float gain) { + PrecisionTensor wt = {.data = t.data, .shape = {rows, cols}}; + puf_kaiming_init(&wt, gain, (*seed)++, stream); + }; + puf_normal_init(&ew->embed_w, 1.0f, (*seed)++, stream); + init2d(ew->proj_w, ew->hidden, N3_CONCAT, 1.0f); + hipMemsetAsync(ew->proj_b.data, 0, numel(ew->proj_b.shape) * sizeof(precision_t), stream); +} + +static void nmmo3_encoder_reg_params(void* w, Allocator* alloc) { + NMMO3EncoderWeights* ew = (NMMO3EncoderWeights*)w; + conv_reg_params(&ew->conv1, alloc); + conv_reg_params(&ew->conv2, alloc); + ew->embed_w = {.shape = {N3_EMBED_VOCAB, N3_EMBED_DIM}}; + ew->proj_w = {.shape = {ew->hidden, N3_CONCAT}}; + ew->proj_b = {.shape = {ew->hidden}}; + alloc_register(alloc,&ew->embed_w); + alloc_register(alloc,&ew->proj_w); alloc_register(alloc,&ew->proj_b); +} + +static void nmmo3_encoder_reg_train(void* w, void* activations, Allocator* acts, Allocator* grads, int B_TT) { + NMMO3EncoderWeights* ew = (NMMO3EncoderWeights*)w; + NMMO3EncoderActivations* a = (NMMO3EncoderActivations*)activations; + *a = {}; + a->multihot = {.shape = {B_TT, N3_MULTIHOT * N3_MAP_H * N3_MAP_W}}; + alloc_register(acts,&a->multihot); + // Conv1 buffers + a->conv1.out = {.shape = {B_TT * N3_C1_OC * N3_C1_OH * N3_C1_OW}}; + a->conv1.grad = {.shape = {B_TT * N3_C1_OC * N3_C1_OH * N3_C1_OW}}; + a->conv1.saved_input = {.shape = {B_TT * N3_C1_IC * N3_MAP_H * N3_MAP_W}}; + a->conv1.wgrad = {.shape = {N3_C1_OC, N3_C1_IC * N3_C1_K * N3_C1_K}}; + a->conv1.bgrad = {.shape = {N3_C1_OC}}; + alloc_register(acts,&a->conv1.out); alloc_register(acts,&a->conv1.grad); alloc_register(acts,&a->conv1.saved_input); + alloc_register(grads,&a->conv1.wgrad); alloc_register(grads,&a->conv1.bgrad); + a->col1 = {.shape = {B_TT * N3_C1_OH * N3_C1_OW, N3_C1_IC * N3_C1_K * N3_C1_K}}; + a->mm1 = {.shape = {B_TT * N3_C1_OH * N3_C1_OW, N3_C1_OC}}; + alloc_register(acts,&a->col1); alloc_register(acts,&a->mm1); + // Conv2 buffers + a->conv2.out = {.shape = {B_TT * N3_C2_OC * N3_C2_OH * N3_C2_OW}}; + a->conv2.grad = {.shape = {B_TT * N3_C2_OC * N3_C2_OH * N3_C2_OW}}; + a->conv2.saved_input = {.shape = {B_TT * N3_C2_IC * N3_C1_OH * N3_C1_OW}}; + a->conv2.wgrad = {.shape = {N3_C2_OC, N3_C2_IC * N3_C2_K * N3_C2_K}}; + a->conv2.bgrad = {.shape = {N3_C2_OC}}; + alloc_register(acts,&a->conv2.out); alloc_register(acts,&a->conv2.grad); alloc_register(acts,&a->conv2.saved_input); + alloc_register(grads,&a->conv2.wgrad); alloc_register(grads,&a->conv2.bgrad); + a->col2 = {.shape = {B_TT * N3_C2_OH * N3_C2_OW, N3_C2_IC * N3_C2_K * N3_C2_K}}; + a->mm2 = {.shape = {B_TT * N3_C2_OH * N3_C2_OW, N3_C2_OC}}; + alloc_register(acts,&a->col2); alloc_register(acts,&a->mm2); + a->embed_out = {.shape = {B_TT, N3_PLAYER_EMBED}}; + a->concat = {.shape = {B_TT, N3_CONCAT}}; + a->out = {.shape = {B_TT, ew->hidden}}; + a->saved_obs = {.shape = {B_TT, ew->obs_size}}; + alloc_register(acts,&a->embed_out); alloc_register(acts,&a->concat); + alloc_register(acts,&a->out); alloc_register(acts,&a->saved_obs); + a->embed_wgrad = {.shape = {N3_EMBED_VOCAB, N3_EMBED_DIM}}; + a->embed_wgrad_f = {.shape = {N3_EMBED_VOCAB, N3_EMBED_DIM}}; + a->proj_wgrad = {.shape = {ew->hidden, N3_CONCAT}}; + a->proj_bgrad = {.shape = {ew->hidden}}; + alloc_register(grads,&a->embed_wgrad); + alloc_register(acts,&a->embed_wgrad_f); + alloc_register(grads,&a->proj_wgrad); alloc_register(grads,&a->proj_bgrad); +} + +static void nmmo3_encoder_reg_rollout(void* w, void* activations, Allocator* alloc, int B) { + NMMO3EncoderWeights* ew = (NMMO3EncoderWeights*)w; + NMMO3EncoderActivations* a = (NMMO3EncoderActivations*)activations; + a->multihot = {.shape = {B, N3_MULTIHOT * N3_MAP_H * N3_MAP_W}}; + alloc_register(alloc,&a->multihot); + a->conv1.out = {.shape = {B * N3_C1_OC * N3_C1_OH * N3_C1_OW}}; + alloc_register(alloc,&a->conv1.out); + a->col1 = {.shape = {B * N3_C1_OH * N3_C1_OW, N3_C1_IC * N3_C1_K * N3_C1_K}}; + a->mm1 = {.shape = {B * N3_C1_OH * N3_C1_OW, N3_C1_OC}}; + alloc_register(alloc,&a->col1); alloc_register(alloc,&a->mm1); + a->conv2.out = {.shape = {B * N3_C2_OC * N3_C2_OH * N3_C2_OW}}; + alloc_register(alloc,&a->conv2.out); + a->col2 = {.shape = {B * N3_C2_OH * N3_C2_OW, N3_C2_IC * N3_C2_K * N3_C2_K}}; + a->mm2 = {.shape = {B * N3_C2_OH * N3_C2_OW, N3_C2_OC}}; + alloc_register(alloc,&a->col2); alloc_register(alloc,&a->mm2); + a->embed_out = {.shape = {B, N3_PLAYER_EMBED}}; + a->concat = {.shape = {B, N3_CONCAT}}; + a->out = {.shape = {B, ew->hidden}}; + alloc_register(alloc,&a->embed_out); alloc_register(alloc,&a->concat); alloc_register(alloc,&a->out); +} + +static void* nmmo3_encoder_create_weights(void* self) { + Encoder* e = (Encoder*)self; + return nmmo3_encoder_create(e->in_dim, e->out_dim); +} +static void nmmo3_encoder_free_weights(void* weights) { free(weights); } +static void nmmo3_encoder_free_activations(void* activations) { free(activations); } + +// Override encoder vtable for known ocean environments. No-op for unknown envs. +static void create_custom_encoder(const std::string& env_name, Encoder* enc) { + if (env_name == "nmmo3") { + *enc = Encoder{ + .forward = nmmo3_encoder_forward, + .backward = nmmo3_encoder_backward, + .init_weights = nmmo3_encoder_init_weights, + .reg_params = nmmo3_encoder_reg_params, + .reg_train = nmmo3_encoder_reg_train, + .reg_rollout = nmmo3_encoder_reg_rollout, + .create_weights = nmmo3_encoder_create_weights, + .free_weights = nmmo3_encoder_free_weights, + .free_activations = nmmo3_encoder_free_activations, + .in_dim = enc->in_dim, .out_dim = enc->out_dim, + .activation_size = sizeof(NMMO3EncoderActivations), + }; + } else if (env_name == "nethack") { + assert(false && "nethack encoder unavailable in HIP build"); + } +} + +static void create_custom_decoder(const std::string& env_name, Decoder* dec) { + if (env_name == "nethack") { assert(false && "nethack decoder unavailable in HIP build"); } +} diff --git a/src-hip/ocean_stub.cpp b/src-hip/ocean_stub.cpp new file mode 100644 index 0000000000..275952bb39 --- /dev/null +++ b/src-hip/ocean_stub.cpp @@ -0,0 +1 @@ +// ocean env GPU encoders are not part of this build diff --git a/src-hip/pufferlib.hip.cpp b/src-hip/pufferlib.hip.cpp new file mode 100644 index 0000000000..db49c780ac --- /dev/null +++ b/src-hip/pufferlib.hip.cpp @@ -0,0 +1,2337 @@ +#include "hip/hip_runtime.h" +#include +#include +// nvtx removed for AMD +#include "nvml.h" +#include +#include + +#include +#include "models.hip.cpp" +#include "ocean.hip.cpp" +#include "muon.hip.cpp" +#include "vecenv.h" +// AMD build: ocean.cu custom encoders are unavailable; stubs live at file end. +static void create_custom_encoder(const std::string& env_name, Encoder* enc); +static void create_custom_decoder(const std::string& env_name, Decoder* dec); + +static double wall_clock() { + struct timespec ts; + clock_gettime(CLOCK_REALTIME, &ts); + return ts.tv_sec + ts.tv_nsec * 1e-9; +} + +enum LossIdx { + LOSS_PG = 0, LOSS_VF = 1, LOSS_ENT = 2, LOSS_TOTAL = 3, + LOSS_OLD_APPROX_KL = 4, LOSS_APPROX_KL = 5, LOSS_CLIPFRAC = 6, + LOSS_N = 7, NUM_LOSSES = 8, +}; + +enum ProfileIdx { + PROF_ROLLOUT = 0, + PROF_EVAL_GPU, + PROF_EVAL_ENV, + PROF_TRAIN_MISC, + PROF_TRAIN_FORWARD, + NUM_PROF, +}; + +static const char* PROF_NAMES[NUM_PROF] = { + "rollout", + "eval_gpu", + "eval_env", + "train_misc", + "train_forward", +}; + +#define NUM_TRAIN_EVENTS 5 +typedef struct { + hipEvent_t events[NUM_TRAIN_EVENTS]; + float accum[NUM_PROF]; +} ProfileT; + +// Data collected by parallel environment workers. Each worker handles +// a constant subset of agents +struct RolloutBuf { + PrecisionTensor observations; // (horizon, agents, input_size) + PrecisionTensor actions; // (horizon, agents, num_atns) + PrecisionTensor values; // (horizon, agents) + PrecisionTensor logprobs; // ... + PrecisionTensor rewards; + PrecisionTensor terminals; + PrecisionTensor ratio; + PrecisionTensor importance; + PrecisionTensor action_mask; // (horizon, agents, mask_size); .data=nullptr when env opts out +}; + +// Buffers are initialized as raw structs with only shape information. alloc_register +// stores the shape and data pointer. Memory is only allocated after all buffers are registered. +void register_rollout_buffers(RolloutBuf& bufs, Allocator* alloc, int T, int B, int input_size, + int num_atns, int mask_size) { + bufs = (RolloutBuf){ + .observations = {.shape = {T, B, input_size}}, + .actions = {.shape = {T, B, num_atns}}, + .values = {.shape = {T, B}}, + .logprobs = {.shape = {T, B}}, + .rewards = {.shape = {T, B}}, + .terminals = {.shape = {T, B}}, + .ratio = {.shape = {T, B}}, + .importance = {.shape = {T, B}}, + .action_mask = {}, + }; + alloc_register(alloc, &bufs.observations); + alloc_register(alloc, &bufs.actions); + alloc_register(alloc, &bufs.values); + alloc_register(alloc, &bufs.logprobs); + alloc_register(alloc, &bufs.rewards); + alloc_register(alloc, &bufs.terminals); + alloc_register(alloc, &bufs.ratio); + alloc_register(alloc, &bufs.importance); + if (mask_size > 0) { + bufs.action_mask = {.shape = {T, B, mask_size}}; + alloc_register(alloc, &bufs.action_mask); + } +} + +// Train data layout is transposed to (B, T) from rollouts layout (T, B) +// This allows env workers to collect data with contiguous writes and +// training to perform several (though not all) ops in contiguous memory +struct TrainGraph { + PrecisionTensor mb_state; // (layers, B, hidden) + PrecisionTensor mb_obs; // (B, T, input_size) + PrecisionTensor mb_actions; // (B, T, num_atns) + PrecisionTensor mb_logprobs; // (B, T) + PrecisionTensor mb_advantages; // ... + PrecisionTensor mb_values; + PrecisionTensor mb_returns; + PrecisionTensor mb_ratio; + PrecisionTensor mb_newvalue; + PrecisionTensor mb_prio; // (B,) + PrecisionTensor mb_action_mask; // (B, T, mask_size); .data=nullptr when disabled +}; + +void register_train_buffers(TrainGraph& bufs, Allocator* alloc, int B, int T, int input_size, + int hidden_size, int num_atns, int num_layers, int mask_size) { + bufs = (TrainGraph){ + .mb_state = {.shape = {num_layers, B, hidden_size}}, + .mb_obs = {.shape = {B, T, input_size}}, + .mb_actions = {.shape = {B, T, num_atns}}, + .mb_logprobs = {.shape = {B, T}}, + .mb_advantages = {.shape = {B, T}}, + .mb_values = {.shape = {B, T}}, + .mb_returns = {.shape = {B, T}}, + .mb_ratio = {.shape = {B, T}}, + .mb_newvalue = {.shape = {B, T}}, + .mb_prio = {.shape = {B}}, + .mb_action_mask = {}, + }; + alloc_register(alloc, &bufs.mb_obs); + alloc_register(alloc, &bufs.mb_state); + alloc_register(alloc, &bufs.mb_actions); + alloc_register(alloc, &bufs.mb_logprobs); + alloc_register(alloc, &bufs.mb_advantages); + alloc_register(alloc, &bufs.mb_prio); + alloc_register(alloc, &bufs.mb_values); + alloc_register(alloc, &bufs.mb_returns); + alloc_register(alloc, &bufs.mb_ratio); + alloc_register(alloc, &bufs.mb_newvalue); + if (mask_size > 0) { + bufs.mb_action_mask = {.shape = {B, T, mask_size}}; + alloc_register(alloc, &bufs.mb_action_mask); + } +} + +// PPO buffers + args are quite complex. We do the entire +// forward + backwards pass for the full loss function in one kernel +struct PPOGraphArgs { + precision_t* out_ratio; + precision_t* out_newvalue; + const precision_t* actions; + const precision_t* old_logprobs; + const precision_t* advantages; + const precision_t* prio; + const precision_t* values; + const precision_t* returns; +}; + +struct PPOKernelArgs { + float* grad_logits; + float* grad_logstd; // For continuous actions + float* grad_values_pred; + const precision_t* logits; + const precision_t* logstd; // Continuous only + const precision_t* values_pred; + const float* adv_mean; + const float* adv_var; + const int* act_sizes; + const precision_t* action_mask; // (N, T, A_total) or nullptr + int mask_stride_n, mask_stride_t; + const signed char* head_consume; // (nverbs, num_atns) or nullptr + int hc_stride; + int num_atns; + float clip_coef, vf_clip_coef, vf_coef; + const float* ent_coef; // device ptr, by-value args get baked into the cuda graph + int T_seq, A_total, N; + int logits_stride_n, logits_stride_t, logits_stride_a; + int values_stride_n, values_stride_t; + bool is_continuous; +}; + +struct PPOBuffersPuf { + FloatTensor loss_output, grad_loss; + FloatTensor saved_for_bwd; + FloatTensor grad_logits, grad_values, grad_logstd, adv_scratch; + FloatTensor ent_coef; +}; + +void register_ppo_buffers(PPOBuffersPuf& bufs, Allocator* alloc, int N, int T, int A_total, bool is_continuous) { + long total = (long)N * T; + bufs = (PPOBuffersPuf){ + .loss_output = {.shape = {1}}, + .grad_loss = {.shape = {1}}, + .saved_for_bwd = {.shape = {total, 5}}, + .grad_logits = {.shape = {N, T, A_total}}, + .grad_values = {.shape = {N, T, 1}}, + .grad_logstd = {.shape = {N, T, A_total}}, + .adv_scratch = {.shape = {2}}, + .ent_coef = {.shape = {1}}, + }; + alloc_register(alloc, &bufs.loss_output); + alloc_register(alloc, &bufs.saved_for_bwd); + alloc_register(alloc, &bufs.grad_loss); + alloc_register(alloc, &bufs.grad_logits); + alloc_register(alloc, &bufs.grad_values); + if (is_continuous) { + alloc_register(alloc, &bufs.grad_logstd); + } + alloc_register(alloc, &bufs.adv_scratch); + alloc_register(alloc, &bufs.ent_coef); +} + +// Prioritized replay over single-epoch data. These kernels are +// the least cleaned because we will likely have a better method in 5.0 +struct PrioBuffers { + FloatTensor prio_probs, cdf, mb_prio; + IntTensor idx; +}; + +void register_prio_buffers(PrioBuffers& bufs, Allocator* alloc, int B, int minibatch_segments) { + bufs = (PrioBuffers){ + .prio_probs = {.shape = {B}}, + .cdf = {.shape = {B}}, + .mb_prio = {.shape = {minibatch_segments}}, + .idx = {.shape = {minibatch_segments}}, + }; + alloc_register(alloc, &bufs.prio_probs); + alloc_register(alloc, &bufs.cdf); + alloc_register(alloc, &bufs.idx); + alloc_register(alloc, &bufs.mb_prio); +} + +// Slice: select dim0 index t, then narrow dim0 from start for count. +// 3D (T, B, F) -> (count, F); 2D (T, B) -> (count,) +inline PrecisionTensor puf_slice(PrecisionTensor& p, int t, int start, int count) { + if (ndim(p.shape) == 3) { + long B = p.shape[1], F = p.shape[2]; + return {.data = p.data + (t*B + start)*F, .shape = {count, F}}; + } else { + long B = p.shape[1]; + return {.data = p.data + (t*B + start), .shape = {count}}; + } +} + +struct EnvBuf { + OBS_TENSOR_T obs; // (total_agents, obs_size) - type defined per-env in binding.c + FloatTensor actions; // (total_agents, num_atns) + FloatTensor rewards; // (total_agents,) + FloatTensor terminals; // (total_agents,) + ByteTensor action_mask; // (total_agents, mask_size); .data=nullptr when env opts out +}; + +StaticVec* create_environments(int num_buffers, int total_agents, + const std::string& env_name, Dict* vec_kwargs, Dict* env_kwargs, EnvBuf& env) { + StaticVec* vec = create_static_vec(total_agents, num_buffers, 1, vec_kwargs, env_kwargs); + env.obs = { + .data = (decltype(env.obs.data))vec->gpu_observations, + .shape = {total_agents, get_obs_size()}, + }; + env.actions = { .data = (float*)vec->gpu_actions, .shape = {total_agents, get_num_atns()} }; + env.rewards = { .data = (float*)vec->gpu_rewards, .shape = {total_agents} }; + env.terminals = { .data = (float*)vec->gpu_terminals, .shape = {total_agents} }; + if (vec->action_mask_size > 0) { + env.action_mask = { .data = vec->gpu_action_mask, + .shape = {total_agents, vec->action_mask_size} }; + } else { + env.action_mask = { .data = nullptr, .shape = {0} }; + } + return vec; +} + +typedef struct { + // Layout + int horizon; + int total_agents; + int num_buffers; + // Model architecture + int num_atns; + int hidden_size; + int num_layers; + // Learning rate + float lr; + float min_lr_ratio; + bool anneal_lr; + // Optimizer + float beta1; + float beta2; + float eps; + // Training + int minibatch_size; + float replay_ratio; + long total_timesteps; + float max_grad_norm; + // PPO + float clip_coef; + float vf_clip_coef; + float vf_coef; + float ent_coef; + // Entropy coefficient anneal — mirrors lr annealing. When anneal_ent_coef + // is set, ent_coef cosine-decays from its base value to + // min_ent_coef_ratio * ent_coef over total_timesteps. + float min_ent_coef_ratio; + bool anneal_ent_coef; + // GAE + float gamma; + float gae_lambda; + // VTrace + float vtrace_rho_clip; + float vtrace_c_clip; + // Priority + float prio_alpha; + float prio_beta0; + // Flags + bool reset_state; + int cudagraphs; + bool profile; + // Multi-GPU + int rank; + int world_size; + int gpu_id; + std::string nccl_id; // raw bytes of ncclUniqueId (empty for single-GPU) + // Threading + int num_threads; + int seed; +} HypersT; + +// A frozen weight bank: same shape as the primary, but its own params buffer +// (and per-buffer rollout states/activations). Used for match (eval) and league +// (frozen historical opponents). Not trained; updated only via load. +typedef struct { + Policy policy; // Bank-owned Policy; lets banks have different arch than primary. + PolicyWeights weights; + Allocator params_alloc; + Allocator acts_alloc; + PrecisionTensor param_puf; + FloatTensor master_weights; + PrecisionTensor* buffer_states; // [num_buffers] + PolicyActivations* buffer_activations; // [num_buffers] + int slice_size; // # agents per buffer this bank owns; sets activation/state batch dim + int hidden_size; + int num_layers; +} WeightBank; + +typedef struct { + Policy policy; + PolicyWeights weights; // current precision_t weights (structured) + PolicyActivations train_activations; + Allocator params_alloc; + Allocator grads_alloc; + Allocator activations_alloc; + StaticVec* vec; + Muon muon; + ncclComm_t nccl_comm; // NCCL communicator for multi-GPU + HypersT hypers; + bool is_continuous; // True if all action dimensions are continuous (size==1) + PrecisionTensor* buffer_states; // Per-buffer states for contiguous access + PolicyActivations* buffer_activations; // Per-buffer inference activations + RolloutBuf rollouts; + RolloutBuf train_rollouts; // Pre-allocated transposed copy for train_impl + EnvBuf env; + TrainGraph train_buf; + PrecisionTensor advantages_puf; // Pre-allocated for train_impl (B, T) + hipGraphExec_t* fused_rollout_cudagraphs; // [horizon][num_buffers] + hipGraphExec_t train_cudagraph; + hipStream_t* streams; // per-buffer raw CUDA streams + hipStream_t default_stream; // main-thread stream (captured once at init) + IntTensor act_sizes_puf; // CUDA int32 tensor of action head sizes + FloatTensor losses_puf; // (NUM_LOSSES,) f32 accumulator + PPOBuffersPuf ppo_bufs_puf; // Pre-allocated buffers for ppo_loss_fwd_bwd + PrioBuffers prio_bufs; // Pre-allocated buffers for prio_replay + FloatTensor master_weights; // fp32 master weights (flat); same buffer as param_puf in fp32 mode + PrecisionTensor param_puf; + PrecisionTensor grad_puf; + LongTensor rng_offset_puf; // (num_buffers+1,) int64 CUDA device counters + ProfileT profile; + nvmlDevice_t nvml_device; + long epoch; + long global_step; + double start_time; + double last_log_time; + long last_log_step; + int train_warmup; + bool rollout_captured; + bool train_captured; + ulong seed; + hiprandStatePhilox4_32_10_t** rng_states; // per-buffer persistent RNG states [num_buffers] + // Optional frozen weight banks for match / league. + WeightBank* frozen_banks; // [num_frozen_banks] + int num_frozen_banks; + std::string env_name; // Kept for post-init bank adds (needs create_custom_encoder). + // Per-buffer-relative bank layout: bank_layout[b] = first agent within each + // buffer chunk owned by bank b. Length num_banks+1; ends at agents_per_buffer. + // Same shape applied to every buffer (each buffer hosts every bank), so each + // worker thread only writes inside its own physical chunk. + // Bank 0 = primary (learner). NULL = no layout set (primary owns full chunk). + int* bank_layout; +} PuffeRL; + +Dict* log_environments_impl(PuffeRL& pufferl) { + // Capacity raised from 32 to 64 to accommodate chess's per-bank + // hist_score_bank_ / hist_n_bank_ entries (16 keys for 8 banks). + Dict* out = create_dict(64); + static_vec_log(pufferl.vec, out); + return out; +} + +inline void profile_begin(const char* tag, bool enable) { + if (enable) nvtxRangePushA(tag); +} + +inline void profile_end(bool enable) { + if (enable) nvtxRangePop(); +} + +// Thread-local stream for per-buffer threads (set once by thread_init_wrapper) +static thread_local hipStream_t tl_stream = 0; + +// Thread initialization callback - sets thread-local stream once per thread +extern "C" void thread_init_wrapper(void* ctx, int buf) { + PuffeRL* pufferl = (PuffeRL*)ctx; + tl_stream = pufferl->streams[buf]; +} + +__global__ void rng_init(hiprandStatePhilox4_32_10_t* states, uint64_t seed, int n) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + hiprand_init(seed, idx, 0, &states[idx]); + } +} + +__device__ __forceinline__ float safe_logit(const precision_t* logits, + int logits_base, int logits_offset, int offset) { + float l = to_float(logits[logits_base + logits_offset + offset]); + if (isnan(l)) { + l = 0.0f; + } + if (isinf(l)) { + l = (l > 0) ? 3.4028e+38f : -3.4028e+38f; + } + return l; +} + +__device__ __forceinline__ float finite_or_clamp(float x, float lo, float hi) { + if (isnan(x)) { + return 0.0f; + } + if (isinf(x)) { + return x > 0.0f ? hi : lo; + } + return fminf(hi, fmaxf(lo, x)); +} + +__device__ __forceinline__ float safe_continuous_mean(const precision_t* logits, int idx) { + return finite_or_clamp(to_float(logits[idx]), -1.0e6f, 1.0e6f); +} + +__device__ __forceinline__ float safe_continuous_logstd(const precision_t* logstd, int idx) { + return finite_or_clamp(to_float(logstd[idx]), -20.0f, 2.0f); +} + +__device__ __forceinline__ float masked_logit(const precision_t* logits, + int logits_base, int logits_offset, int offset, + const precision_t* mask, int mask_base) { + float l = safe_logit(logits, logits_base, logits_offset, offset); + if (mask != nullptr) { + float m = to_float(mask[mask_base + logits_offset + offset]); + if (m == 0.0f) l = -1e4f; + } + return l; +} + +// Expects action logits and values to be in the same contiguous buffer. See default decoder +// ---- consumed-head gating (opt-in via PUFFER_HEAD_GATING) ---- +extern "C" __attribute__((weak)) const signed char* env_head_consume_map(int*, int*); +static const signed char* g_hc_dev = nullptr; +static int g_hc_stride = 0; +static bool g_hc_init = false; +static const signed char* get_head_consume_dev(int* stride) { + if (!g_hc_init) { + g_hc_init = true; + const char* hg = getenv("PUFFER_HEAD_GATING"); + if (hg && hg[0] && hg[0] != '0' && env_head_consume_map) { + int nv = 0, na = 0; + const signed char* host = env_head_consume_map(&nv, &na); + if (host && nv > 0 && na > 0) { + signed char* dev = nullptr; + hipMalloc(&dev, (size_t)nv * na); + hipMemcpy(dev, host, (size_t)nv * na, hipMemcpyHostToDevice); + g_hc_dev = dev; g_hc_stride = na; + } + } + } + *stride = g_hc_stride; + return g_hc_dev; +} + +__global__ void sample_logits( + PrecisionTensor dec_out, // (B, logits_dim + 1 for values) + PrecisionTensor logstd_puf, // (1, od) - continuous actions only + IntTensor act_sizes_puf, // (num_atns,) action head sizes + precision_t* __restrict__ actions, // (B, num_atns) + precision_t* __restrict__ logprobs, // (B,) + precision_t* __restrict__ value_out, // (B,) + hiprandStatePhilox4_32_10_t* __restrict__ rng_states, + const precision_t* __restrict__ action_mask, // (B, A_total) or nullptr + int mask_stride, // 0 when action_mask is nullptr + const signed char* __restrict__ head_consume, // (nverbs, num_atns) or nullptr + int hc_stride) { + int B = dec_out.shape[0]; + int fused_cols = dec_out.shape[1]; + int num_atns = numel(act_sizes_puf.shape); + const int* act_sizes = act_sizes_puf.data; + const precision_t* logits = dec_out.data; + int logits_stride = fused_cols; + int value_stride = fused_cols; + bool is_continuous = logstd_puf.data != nullptr && numel(logstd_puf.shape) > 0; + const precision_t* logstd = logstd_puf.data; + int logstd_stride = is_continuous ? 0 : 0; // 1D broadcast: stride 0 + const precision_t* value = logits + (fused_cols - 1); // last column + + int idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= B) { + return; + } + + // Load persistent RNG state (advanced in-place each call) + hiprandStatePhilox4_32_10_t state = rng_states[idx]; + + int logits_base = idx * logits_stride; + float total_log_prob = 0.0f; + + if (is_continuous) { + // Continuous action sampling from Normal(mean, exp(logstd)) + constexpr float LOG_2PI = 1.8378770664093453f; // log(2*pi) + int logstd_base = idx * logstd_stride; // separate stride for logstd (may be 0 for broadcast) + + for (int h = 0; h < num_atns; ++h) { + float mean = safe_continuous_mean(logits, logits_base + h); + float log_std = safe_continuous_logstd(logstd, logstd_base + h); + float std = expf(log_std); + + // Sample from N(0,1) and transform: action = mean + std * noise + float noise = hiprand_normal(&state); + float action = finite_or_clamp(mean + std * noise, -1.0e6f, 1.0e6f); + + precision_t stored_action_p = from_float(action); + float stored_action = to_float(stored_action_p); + // Log probability: -0.5 * ((action - mean) / std)^2 - 0.5 * log(2*pi) - log(std) + float normalized = (stored_action - mean) / std; + float log_prob = -0.5f * normalized * normalized - 0.5f * LOG_2PI - log_std; + + actions[idx * num_atns + h] = stored_action_p; + total_log_prob += log_prob; + } + } else { + // Discrete action sampling (original multinomial logic) + int logits_offset = 0; // offset within row for current action head + int mask_base = (action_mask != nullptr) ? idx * mask_stride : 0; + + for (int h = 0; h < num_atns; ++h) { + int A = act_sizes[h]; // size of this action head + + // Step 1: Find max and sum for numerical stability (with nan_to_num) + float max_val = -INFINITY; + float sum_exp = 0.0f; + for (int a = 0; a < A; ++a) { + float l = masked_logit(logits, logits_base, logits_offset, a, action_mask, mask_base); + if (l > max_val) { + sum_exp *= expf(max_val - l); + max_val = l; + } + sum_exp += expf(l - max_val); + } + float logsumexp = max_val + logf(sum_exp); + + // Step 3: Generate random value for this action head + float rand_val = hiprand_uniform(&state); + + // Step 4: Multinomial sampling using inverse CDF + float cumsum = 0.0f; + int sampled_action = -1; // sentinel: no action chosen yet + + for (int a = 0; a < A; ++a) { + float l = masked_logit(logits, logits_base, logits_offset, a, action_mask, mask_base); + float prob = expf(l - logsumexp); + cumsum += prob; + if (rand_val < cumsum) { + sampled_action = a; + break; + } + } + + // Float rounding can leave cumsum < 1.0; fall back to the last legal action. + if (sampled_action < 0) { + sampled_action = A - 1; + if (action_mask != nullptr) { + for (int a = A - 1; a >= 0; --a) { + if (to_float(action_mask[mask_base + logits_offset + a]) != 0.0f) { + sampled_action = a; + break; + } + } + } + } + + // Step 5: Gather log probability of sampled action + float sampled_logit = masked_logit(logits, logits_base, logits_offset, sampled_action, action_mask, mask_base); + float log_prob = sampled_logit - logsumexp; + + // Write action for this head + actions[idx * num_atns + h] = from_float(sampled_action); + // consumed-head gating: only heads the sampled verb uses count + int verb = (int)to_float(actions[idx * num_atns]); + int used = (head_consume == nullptr || h == 0) + ? 1 : (int)head_consume[verb * hc_stride + h]; + if (used) total_log_prob += log_prob; + + // Advance to next action head + logits_offset += A; + } + } + + // Write summed log probability (log of joint probability) + logprobs[idx] = from_float(total_log_prob); + + // Copy value (fused to avoid separate elementwise kernel for strided->contiguous copy) + value_out[idx] = value[idx * value_stride]; + + // Save RNG state back for next call + rng_states[idx] = state; +} + +// Single step rollout forward pass. Called by each environment worker in their +// own buffer thread. This operation is cudagraphed. +extern "C" void net_callback_wrapper(void* ctx, int buf, int t) { + PuffeRL* pufferl = (PuffeRL*)ctx; + HypersT& hypers = pufferl->hypers; + int graph = t * hypers.num_buffers + buf; + profile_begin("fused_rollout", hypers.profile); + + hipStream_t current_stream = tl_stream; + if (pufferl->rollout_captured) { + assert(hipGraphLaunch(pufferl->fused_rollout_cudagraphs[graph], current_stream) == hipSuccess + && "hipGraphLaunch failed"); + profile_end(hypers.profile); + return; + } + + bool capturing = pufferl->epoch == hypers.cudagraphs; + if (capturing) { + assert(hipStreamBeginCapture(current_stream, hipStreamCaptureModeGlobal) == hipSuccess + && "hipStreamBeginCapture failed"); + } + + RolloutBuf& rollouts = pufferl->rollouts; + EnvBuf& env = pufferl->env; + int block_size = pufferl->vec->total_agents / hypers.num_buffers; + int start = buf * block_size; + hipStream_t stream = current_stream; + + // Copy observations, rewards, terminals from GPU env buffers to rollout buffer + OBS_TENSOR_T& obs_env = env.obs; + int n = block_size * obs_env.shape[1]; + PrecisionTensor obs_dst = puf_slice(rollouts.observations, t, start, block_size); + cast_dispatch(obs_dst.data, obs_env.data + (long)start*obs_env.shape[1], n, stream); + + PrecisionTensor rew_dst = puf_slice(rollouts.rewards, t, start, block_size); + n = block_size; + cast<<>>( + rew_dst.data, env.rewards.data + start, n); + + PrecisionTensor term_dst = puf_slice(rollouts.terminals, t, start, block_size); + cast<<>>( + term_dst.data, env.terminals.data + start, n); + + // Copy action mask from env into rollout buffer (if env opted in) + PrecisionTensor mask_slice = {}; + int mask_stride = 0; + if (rollouts.action_mask.data != nullptr) { + int mask_size = rollouts.action_mask.shape[2]; + mask_stride = mask_size; + mask_slice = puf_slice(rollouts.action_mask, t, start, block_size); + int mask_n = block_size * mask_size; + cast<<>>( + mask_slice.data, + env.action_mask.data + (long)start * mask_size, + mask_n); + } + + // Per-bank policy forward + sampling. Each bank owns a contiguous sub-range + // [bank_layout[b], bank_layout[b+1]) within every buffer's chunk; layout is + // per-buffer-relative so each worker writes only inside its own chunk. + // Cudagraph capture absorbs the extra kernel launches. + int num_banks = 1 + pufferl->num_frozen_banks; + long act_cols = env.actions.shape[1]; + for (int b = 0; b < num_banks; b++) { + int bank_off = pufferl->bank_layout ? pufferl->bank_layout[b] : 0; + int bank_end = pufferl->bank_layout ? pufferl->bank_layout[b + 1] : block_size; + int bank_size = bank_end - bank_off; + if (bank_size == 0) continue; + + Policy* p_bank; + PolicyWeights* w_bank; + PolicyActivations* a_bank; + PrecisionTensor* s_bank; + if (b == 0) { + p_bank = &pufferl->policy; + w_bank = &pufferl->weights; + a_bank = &pufferl->buffer_activations[buf]; + s_bank = &pufferl->buffer_states[buf]; + } else { + WeightBank* fb = &pufferl->frozen_banks[b - 1]; + p_bank = &fb->policy; + w_bank = &fb->weights; + a_bank = &fb->buffer_activations[buf]; + s_bank = &fb->buffer_states[buf]; + } + + int sub_start = start + bank_off; + PrecisionTensor obs_b = puf_slice(rollouts.observations, t, sub_start, bank_size); + PrecisionTensor act_b = puf_slice(rollouts.actions, t, sub_start, bank_size); + PrecisionTensor lp_b = puf_slice(rollouts.logprobs, t, sub_start, bank_size); + PrecisionTensor val_b = puf_slice(rollouts.values, t, sub_start, bank_size); + PrecisionTensor mask_b = {}; + int mask_stride_b = 0; + if (rollouts.action_mask.data != nullptr) { + mask_b = puf_slice(rollouts.action_mask, t, sub_start, bank_size); + mask_stride_b = mask_stride; + } + + PrecisionTensor dec_puf = policy_forward(p_bank, *w_bank, *a_bank, obs_b, *s_bank, stream); + + PrecisionTensor p_logstd = {}; + DecoderWeights* dw = (DecoderWeights*)w_bank->decoder; + if (dw->continuous) { + p_logstd = dw->logstd; + } + + // Offset RNG by bank_off so banks don't collide on per-buffer rng slots. + int hc_stride_s = 0; + const signed char* hc_dev_s = get_head_consume_dev(&hc_stride_s); + sample_logits<<>>( + dec_puf, p_logstd, pufferl->act_sizes_puf, + act_b.data, lp_b.data, val_b.data, + pufferl->rng_states[buf] + bank_off, + mask_b.data, mask_stride_b, hc_dev_s, hc_stride_s); + + cast<<>>( + env.actions.data + (long)sub_start * act_cols, + act_b.data, numel(act_b.shape)); + } + + if (capturing) { + hipGraph_t _graph; + assert(hipStreamEndCapture(current_stream, &_graph) == hipSuccess + && "hipStreamEndCapture failed"); + assert(hipGraphInstantiateWithFlags(&pufferl->fused_rollout_cudagraphs[graph], _graph, 0) == hipSuccess + && "hipGraphInstantiate failed"); + assert(hipGraphDestroy(_graph) == hipSuccess && "hipGraphDestroy failed"); + hipDeviceSynchronize(); + } + profile_end(hypers.profile); +} + + +__device__ __forceinline__ float load_logit_masked( + const precision_t* __restrict__ logits, int logits_base, + int logits_stride_a, int logits_offset, int a, + const precision_t* __restrict__ mask, int mask_base) { + float l = to_float(logits[logits_base + (logits_offset + a) * logits_stride_a]); + if (mask != nullptr) { + float m = to_float(mask[mask_base + logits_offset + a]); + if (m == 0.0f) { + l = -1e4f; + return l; + } + } + return l; +} + +__device__ __forceinline__ void ppo_discrete_head( + const precision_t* __restrict__ logits, int logits_base, + int logits_stride_a, int logits_offset, int A, int act, + const precision_t* __restrict__ mask, int mask_base, + float* out_logsumexp, float* out_entropy, float* out_logp) { + float max_logit = -INFINITY; + float sum = 0.0f; + float act_logit = 0.0f; + + for (int a = 0; a < A; ++a) { + float l = load_logit_masked(logits, logits_base, logits_stride_a, logits_offset, a, mask, mask_base); + if (a == act) { + act_logit = l; + } + if (l > max_logit) { + sum *= __expf(max_logit - l); + max_logit = l; + } + sum += __expf(l - max_logit); + } + float logsumexp = max_logit + __logf(sum); + + float ent = 0.0f; + for (int a = 0; a < A; ++a) { + float l = load_logit_masked(logits, logits_base, logits_stride_a, logits_offset, a, mask, mask_base); + float logp = l - logsumexp; + float p = __expf(logp); + ent -= p * logp; + } + + *out_logsumexp = logsumexp; + *out_entropy = ent; + *out_logp = act_logit - logsumexp; +} + +__device__ __forceinline__ void ppo_continuous_head( + float mean, float log_std, float action, + float* out_logp, float* out_entropy) { + constexpr float HALF_LOG_2PI = 0.9189385332046727f; + constexpr float HALF_1_PLUS_LOG_2PI = 1.4189385332046727f; + float std = __expf(log_std); + float normalized = (action - mean) / std; + *out_logp = -0.5f * normalized * normalized - HALF_LOG_2PI - log_std; + *out_entropy = HALF_1_PLUS_LOG_2PI + log_std; +} + +__global__ void ppo_loss_compute( + float* __restrict__ ppo_partials, + PPOKernelArgs a, PPOGraphArgs g) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int tid = threadIdx.x; + int total_elements = a.N * a.T_seq; + float inv_NT = 1.0f / float(total_elements); + + __shared__ float block_losses[LOSS_N][PPO_THREADS]; + for (int c = 0; c < LOSS_N; c++) { + block_losses[c][tid] = 0.0f; + } + + if (idx >= total_elements) { + goto reduce; + } + + { + int n = idx / a.T_seq; + int t = idx % a.T_seq; + int nt = n * a.T_seq + t; + + int logits_base = n * a.logits_stride_n + t * a.logits_stride_t; + int values_idx = n * a.values_stride_n + t * a.values_stride_t; + int grad_logits_base = nt * a.A_total; + + // Shared computation (used by both forward and backward) + + float old_logp = to_float(g.old_logprobs[nt]); + float adv = to_float(g.advantages[nt]); + float w = to_float(g.prio[n]); + float val = to_float(g.values[nt]); + float ret = to_float(g.returns[nt]); + float val_pred = to_float(a.values_pred[values_idx]); + g.out_newvalue[nt] = from_float(val_pred); + + float adv_std = sqrtf(float(a.adv_var[0])); + float adv_normalized = (adv - float(a.adv_mean[0])) / (adv_std + 1e-8f); + + // grad_loss is always 1.0 (set in post_create, never changes) + float dL = inv_NT; + float d_pg_loss = dL; + float ent_coef = *a.ent_coef; + float d_entropy_term = dL * (-ent_coef); + + // Value loss (forward) + value gradient (backward) + + float v_error = val_pred - val; + float v_clipped = val + fmaxf(-a.vf_clip_coef, fminf(a.vf_clip_coef, v_error)); + float v_loss_unclipped = (val_pred - ret) * (val_pred - ret); + float v_loss_clipped = (v_clipped - ret) * (v_clipped - ret); + float v_loss = 0.5f * fmaxf(v_loss_unclipped, v_loss_clipped); + + // Value gradient + bool use_clipped_vf = (v_loss_clipped > v_loss_unclipped); + float d_val_pred = 0.0f; + if (use_clipped_vf) { + if (v_error >= -a.vf_clip_coef && v_error <= a.vf_clip_coef) { + d_val_pred = v_clipped - ret; + } + } else { + d_val_pred = val_pred - ret; + } + a.grad_values_pred[nt] = dL * a.vf_coef * d_val_pred; + + // Policy loss + gradients + + float pg_loss, total_entropy, logratio, ratio; + float total_log_prob = 0.0f; + total_entropy = 0.0f; + + // Discrete-only: per-head arrays needed across forward + backward + float head_logsumexp[MAX_ATN_HEADS]; + float head_entropy[MAX_ATN_HEADS]; + int head_act[MAX_ATN_HEADS]; + int head_used[MAX_ATN_HEADS]; + + int mask_base = (a.action_mask != nullptr) + ? n * a.mask_stride_n + t * a.mask_stride_t : 0; + + if (!a.is_continuous) { + // consumed-head gating: heads the sampled verb (head 0) does not use + // contribute no logprob/entropy/gradient (see env_head_consume_map) + int verb = static_cast(g.actions[nt * a.num_atns]); + int logits_offset = 0; + for (int h = 0; h < a.num_atns; ++h) { + int A = a.act_sizes[h]; + int act = static_cast(g.actions[nt * a.num_atns + h]); + head_act[h] = act; + int used = (a.head_consume == nullptr || h == 0) + ? 1 : (int)a.head_consume[verb * a.hc_stride + h]; + head_used[h] = used; + float lse, ent, lp; + ppo_discrete_head(a.logits, logits_base, a.logits_stride_a, logits_offset, A, act, + a.action_mask, mask_base, &lse, &ent, &lp); + head_logsumexp[h] = lse; + head_entropy[h] = ent; + if (used) { total_log_prob += lp; total_entropy += ent; } + logits_offset += A; + } + } else { + for (int h = 0; h < a.num_atns; ++h) { + float mean = safe_continuous_mean(a.logits, logits_base + h * a.logits_stride_a); + float log_std = safe_continuous_logstd(a.logstd, h); + float action = finite_or_clamp(float(g.actions[nt * a.num_atns + h]), -1.0e6f, 1.0e6f); + float lp, ent; + ppo_continuous_head(mean, log_std, action, &lp, &ent); + total_log_prob += lp; + total_entropy += ent; + } + } + + // Shared pg loss computation + logratio = total_log_prob - old_logp; + ratio = __expf(logratio); + g.out_ratio[nt] = from_float(ratio); + float ratio_clipped = fmaxf(1.0f - a.clip_coef, fminf(1.0f + a.clip_coef, ratio)); + float wa = -w * adv_normalized; + float pg_loss1 = wa * ratio; + float pg_loss2 = wa * ratio_clipped; + pg_loss = fmaxf(pg_loss1, pg_loss2); + + float d_ratio = wa * d_pg_loss; + if (pg_loss2 > pg_loss1) { + if (ratio <= (1.0f - a.clip_coef) || ratio >= (1.0f + a.clip_coef)) { + d_ratio = 0.0f; + } + } + float d_new_logp = d_ratio * ratio; + + if (!a.is_continuous) { + int logits_offset = 0; + for (int h = 0; h < a.num_atns; ++h) { + int A = a.act_sizes[h]; + if (!head_used[h]) { // gated head: no gradient + for (int j = 0; j < A; ++j) + a.grad_logits[grad_logits_base + logits_offset + j] = 0.0f; + logits_offset += A; + continue; + } + int act = head_act[h]; + float logsumexp = head_logsumexp[h]; + float ent = head_entropy[h]; + + for (int j = 0; j < A; ++j) { + float l = load_logit_masked(a.logits, logits_base, a.logits_stride_a, + logits_offset, j, a.action_mask, mask_base); + float logp = l - logsumexp; + float p = __expf(logp); + float d_logit = (j == act) ? d_new_logp : 0.0f; + d_logit -= p * d_new_logp; + d_logit += d_entropy_term * p * (-ent - logp); + a.grad_logits[grad_logits_base + logits_offset + j] = d_logit; + } + logits_offset += A; + } + } else { + for (int h = 0; h < a.num_atns; ++h) { + float mean = safe_continuous_mean(a.logits, logits_base + h * a.logits_stride_a); + float log_std = safe_continuous_logstd(a.logstd, h); + float std = __expf(log_std); + float var = std * std; + float action = finite_or_clamp(float(g.actions[nt * a.num_atns + h]), -1.0e6f, 1.0e6f); + float diff = action - mean; + + a.grad_logits[grad_logits_base + h] = d_new_logp * diff / var; + a.grad_logstd[nt * a.num_atns + h] = d_new_logp * (diff * diff / var - 1.0f) + d_entropy_term; + } + } + + // Forward: loss partials + float thread_loss = (pg_loss + a.vf_coef * v_loss - ent_coef * total_entropy) * inv_NT; + block_losses[LOSS_PG][tid] = pg_loss * inv_NT; + block_losses[LOSS_VF][tid] = v_loss * inv_NT; + block_losses[LOSS_ENT][tid] = total_entropy * inv_NT; + block_losses[LOSS_TOTAL][tid] = thread_loss; + block_losses[LOSS_OLD_APPROX_KL][tid] = (-logratio) * inv_NT; + block_losses[LOSS_APPROX_KL][tid] = ((ratio - 1.0f) - logratio) * inv_NT; + block_losses[LOSS_CLIPFRAC][tid] = (fabsf(ratio - 1.0f) > a.clip_coef ? 1.0f : 0.0f) * inv_NT; + } // end if (idx < total_elements) + +// Deterministic aggregation +reduce: + __syncthreads(); + + for (int stride = PPO_THREADS / 2; stride > 0; stride >>= 1) { + if (tid < stride) { + for (int c = 0; c < LOSS_N; c++) { + block_losses[c][tid] += block_losses[c][tid + stride]; + } + } + __syncthreads(); + } + + if (tid == 0) { + int base = blockIdx.x * (LOSS_N + 1); + ppo_partials[base] = block_losses[LOSS_TOTAL][0]; + for (int c = 0; c < LOSS_N; c++) { + ppo_partials[base + 1 + c] = block_losses[c][0]; + } + } +} + +// Deterministic reduction of per-block PPO loss partials + count increment +__global__ void ppo_loss_reduce( + float* __restrict__ loss, + float* __restrict__ losses_acc, + const float* __restrict__ partials, + int num_blocks) { + int tid = threadIdx.x; + if (tid > LOSS_N) { + return; + } + + float sum = 0.0f; + for (int b = 0; b < num_blocks; b++) { + sum += partials[b * (LOSS_N + 1) + tid]; + } + + if (tid == 0) { + *loss += sum; + } else { + losses_acc[tid - 1] += sum; + } + + // Fold add_scalar: increment epoch count + if (tid == 0) { + losses_acc[LOSS_N] += 1.0f; + } +} + +__global__ void ppo_var_mean(const precision_t* __restrict__ src, + float* __restrict__ var_out, float* __restrict__ mean_out, int n) { + __shared__ float sdata[256]; + int tid = threadIdx.x; + float sum = 0.0f; + for (int i = tid; i < n; i += blockDim.x) { + sum += to_float(src[i]); + } + sdata[tid] = sum; + __syncthreads(); + for (int s = blockDim.x / 2; s > 0; s >>= 1) { + if (tid < s) { + sdata[tid] += sdata[tid + s]; + } + __syncthreads(); + } + float mean = sdata[0] / (float)n; + if (tid == 0) { + *mean_out = mean; + } + __syncthreads(); + float ss = 0.0f; + for (int i = tid; i < n; i += blockDim.x) { + float d = to_float(src[i]) - mean; + ss += d * d; + } + sdata[tid] = ss; + __syncthreads(); + for (int s = blockDim.x / 2; s > 0; s >>= 1) { + if (tid < s) { + sdata[tid] += sdata[tid + s]; + } + __syncthreads(); + } + if (tid == 0) { + *var_out = sdata[0] / (float)(n - 1); + } +} + +// This is a huge kernel for a relatively cheap operation. But without this, +// it's death by a thousand cuts with repeated kernel launches. Even graphed, you +// blow up the memory bandwidth. +void ppo_loss_fwd_bwd( + PrecisionTensor& dec_out, // (N, T, fused_cols) — fused logits+value from decoder + PrecisionTensor& logstd, // continuous logstd or empty + TrainGraph& graph, + IntTensor& act_sizes, FloatTensor& losses_acc, + float clip_coef, float vf_clip_coef, float vf_coef, const float* ent_coef, + PPOBuffersPuf& bufs, bool is_continuous, + hipStream_t stream) { + int N = dec_out.shape[0], T = dec_out.shape[1], fused_cols = dec_out.shape[2]; + int A_total = fused_cols - 1; // last column is value + int total = N * T; + + // Pointers into fused decoder output + const precision_t* logits_ptr = dec_out.data; + + float* adv_var_ptr = bufs.adv_scratch.data; + float* adv_mean_ptr = adv_var_ptr + 1; + ppo_var_mean<<<1, 256, 0, stream>>>( + graph.mb_advantages.data, adv_var_ptr, adv_mean_ptr, numel(graph.mb_advantages.shape)); + + int ppo_grid = (total + PPO_THREADS - 1) / PPO_THREADS; + + static float* ppo_partials_buf = nullptr; + static int ppo_partials_capacity = 0; + int ppo_partials_needed = ppo_grid * (LOSS_N + 1); + if (!ppo_partials_buf || ppo_partials_needed > ppo_partials_capacity) { + if (ppo_partials_buf) hipFree(ppo_partials_buf); + ppo_partials_capacity = ppo_partials_needed; + hipMalloc(&ppo_partials_buf, ppo_partials_capacity * sizeof(float)); + } + + hipMemsetAsync(bufs.loss_output.data, 0, sizeof(float), stream); + + PPOGraphArgs graph_args = { + .out_ratio = graph.mb_ratio.data, + .out_newvalue = graph.mb_newvalue.data, + .actions = graph.mb_actions.data, + .old_logprobs = graph.mb_logprobs.data, + .advantages = graph.mb_advantages.data, + .prio = graph.mb_prio.data, + .values = graph.mb_values.data, + .returns = graph.mb_returns.data, + }; + + bool has_mask = (graph.mb_action_mask.data != nullptr); + int hc_stride_l = 0; + const signed char* hc_dev_l = get_head_consume_dev(&hc_stride_l); +PPOKernelArgs args = { + .grad_logits = bufs.grad_logits.data, + .grad_logstd = is_continuous ? bufs.grad_logstd.data : nullptr, + .grad_values_pred = bufs.grad_values.data, + .logits = logits_ptr, + .logstd = is_continuous ? logstd.data : nullptr, + .values_pred = logits_ptr + A_total, + .adv_mean = adv_mean_ptr, + .adv_var = adv_var_ptr, + .act_sizes = act_sizes.data, + .action_mask = has_mask ? graph.mb_action_mask.data : nullptr, + .mask_stride_n = has_mask ? T * A_total : 0, + .mask_stride_t = has_mask ? A_total : 0, + .head_consume = hc_dev_l, + .hc_stride = hc_stride_l, + .num_atns = (int)numel(act_sizes.shape), + .clip_coef = clip_coef, .vf_clip_coef = vf_clip_coef, + .vf_coef = vf_coef, .ent_coef = ent_coef, + .T_seq = T, .A_total = A_total, .N = N, + .logits_stride_n = T * fused_cols, .logits_stride_t = fused_cols, .logits_stride_a = 1, + .values_stride_n = T * fused_cols, .values_stride_t = fused_cols, + .is_continuous = is_continuous, + }; + + ppo_loss_compute<<>>(ppo_partials_buf, args, graph_args); + + ppo_loss_reduce<<<1, LOSS_N + 1, 0, stream>>>( + bufs.loss_output.data, losses_acc.data, ppo_partials_buf, ppo_grid); +} + +#define PRIO_WARP_SIZE 32 +#define PRIO_FULL_MASK 0xffffffffull +#define PRIO_BLOCK_SIZE 256 +#define PRIO_NUM_WARPS (PRIO_BLOCK_SIZE / PRIO_WARP_SIZE) +__global__ void compute_prio_adv_reduction( + const precision_t* __restrict__ advantages, + float* prio_weights, float prio_alpha, int stride) { + int row = blockIdx.x; + int tx = threadIdx.x; + int offset = row * stride; + + float local_sum = 0.0f; + for (int t = tx; t < stride; t += blockDim.x) { + local_sum += fabsf(to_float(advantages[offset + t])); + } + + for (int s = PRIO_WARP_SIZE / 2; s >= 1; s /= 2) { + local_sum += __shfl_down_sync(PRIO_FULL_MASK, local_sum, s); + } + if (tx == 0) { + float pw = __powf(local_sum, prio_alpha); + if (isnan(pw) || isinf(pw)) { + pw = 0.0f; + } + prio_weights[row] = pw; + } +} + +__global__ void compute_prio_normalize(float* prio_weights, int length) { + __shared__ float shmem[PRIO_NUM_WARPS]; + __shared__ float block_sum; + + int tx = threadIdx.x; + int lane = tx % PRIO_WARP_SIZE; + int warp_id = tx / PRIO_WARP_SIZE; + const float eps = 1e-6f; + + float local_sum = 0.0f; + for (int t = tx; t < length; t += blockDim.x) { + local_sum += prio_weights[t]; + } + for (int s = PRIO_WARP_SIZE / 2; s >= 1; s /= 2) { + local_sum += __shfl_down_sync(PRIO_FULL_MASK, local_sum, s); + } + if (lane == 0) { + shmem[warp_id] = local_sum; + } + __syncthreads(); + + if (warp_id == 0) { + float val = (lane < PRIO_NUM_WARPS) ? shmem[lane] : 0.0f; + for (int s = PRIO_NUM_WARPS / 2; s >= 1; s /= 2) { + val += __shfl_down_sync(PRIO_FULL_MASK, val, s); + } + if (tx == 0) { + block_sum = val + eps; + } + } + __syncthreads(); + + for (int t = tx; t < length; t += blockDim.x) { + prio_weights[t] = (prio_weights[t] + eps) / block_sum; + } +} + +// mb_prio[i] = pow(total_agents * prio_probs[idx[i]], -anneal_beta) +__global__ void compute_prio_imp_weights( + const int* __restrict__ indices, + const float* __restrict__ prio_probs, + float* mb_prio, int total_agents, + float anneal_beta, int minibatch_segments) { + int tx = threadIdx.x + blockIdx.x * blockDim.x; + if (tx < minibatch_segments) { + float value = prio_probs[indices[tx]] * (float)total_agents; + mb_prio[tx] = __powf(value, -anneal_beta); + } +} + +__global__ void build_cdf( + float* __restrict__ cdf, const float* __restrict__ probs, int B) { + if (blockIdx.x == 0 && threadIdx.x == 0) { + float cum = 0.0f; + for (int i = 0; i < B; i++) { + cum += probs[i]; + cdf[i] = cum; + } + } +} + +__global__ void advance_rng_offset(int64_t* __restrict__ offset_ptr, int64_t delta) { + if (blockIdx.x == 0 && threadIdx.x == 0) { + *offset_ptr += delta; + } +} + +// Multinomial with replacement (uses cuRAND) +__global__ void multinomial_sample(int* __restrict__ out_idx, const float* __restrict__ cdf, + int B, int num_samples, uint64_t seed, const int64_t* __restrict__ offset_ptr) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= num_samples) return; + + uint64_t base_off = (uint64_t)(*offset_ptr); + hiprandStatePhilox4_32_10_t rng_state; + hiprand_init(seed, base_off + tid, 0, &rng_state); + float u = hiprand_uniform(&rng_state); + + int lo = 0, hi = B - 1; + while (lo < hi) { + int mid = (lo + hi) / 2; + if (cdf[mid] < u) lo = mid + 1; + else hi = mid; + } + out_idx[tid] = lo; +} + +// Prioritize high absolute advantage trajectories +// This is a form of implicit curriculum learning +// It is a major improvement in some complex environments +// The values of alpha and beta found by sweeps will tell you +// whether it is important for your task +void prio_replay_cuda(PrecisionTensor& advantages, float prio_alpha, + int minibatch_segments, int total_agents, float anneal_beta, + PrioBuffers& bufs, ulong seed, long* offset_ptr, hipStream_t stream) { + int B = advantages.shape[0], T = advantages.shape[1]; + compute_prio_adv_reduction<<>>( + advantages.data, bufs.prio_probs.data, prio_alpha, T); + compute_prio_normalize<<<1, PRIO_BLOCK_SIZE, 0, stream>>>( + bufs.prio_probs.data, B); + //int block = fmaxf(((minibatch_segments + 31) / 32) * 32, 32); + build_cdf<<<1, 1, 0, stream>>>(bufs.cdf.data, bufs.prio_probs.data, B); + int threads = 256; + int blocks = (minibatch_segments + threads - 1) / threads; + multinomial_sample<<>>( + bufs.idx.data, bufs.cdf.data, B, minibatch_segments, seed, offset_ptr); + advance_rng_offset<<<1, 1, 0, stream>>>(offset_ptr, (int64_t)minibatch_segments); + + int p3_blocks = (minibatch_segments + PRIO_BLOCK_SIZE - 1) / PRIO_BLOCK_SIZE; + compute_prio_imp_weights<<>>( + bufs.idx.data, bufs.prio_probs.data, + bufs.mb_prio.data, total_agents, anneal_beta, minibatch_segments); +} + +// Experience the puffer advantage! Generalized advantage estimation + V-Trace +// importance sampling correction in a single streamlined operation +__device__ void puff_advantage_row_scalar( + const precision_t* values, const precision_t* rewards, const precision_t* dones, + const precision_t* importance, precision_t* advantages, float gamma, float lambda, + float rho_clip, float c_clip, int horizon) { + float lastpufferlam = 0; + for (int t = horizon-2; t >= 0; t--) { + int t_next = t + 1; + float nextnonterminal = 1.0f - to_float(dones[t_next]); + float imp = to_float(importance[t]); + float rho_t = fminf(imp, rho_clip); + float c_t = fminf(imp, c_clip); + float r_nxt = to_float(rewards[t_next]); + float v = to_float(values[t]); + float v_nxt = to_float(values[t_next]); + float delta = rho_t*r_nxt + gamma*v_nxt*nextnonterminal - v; + lastpufferlam = delta + gamma*lambda*c_t*lastpufferlam*nextnonterminal; + advantages[t] = from_float(lastpufferlam); + } +} + +// These loading fns just optimize bandwidth for advantage since we call it on all +// the data every minibatch. This should change in 5.0 +__device__ __forceinline__ void adv_vec_load(const float* ptr, float* out) { + float4 v = *reinterpret_cast(ptr); + out[0] = v.x; out[1] = v.y; out[2] = v.z; out[3] = v.w; +} + +__device__ __forceinline__ void adv_vec_load(const __hip_bfloat16* ptr, float* out) { + uint4 raw = *reinterpret_cast(ptr); + const __hip_bfloat16* bf = reinterpret_cast(&raw); + #pragma unroll + for (int i = 0; i < 8; i++) { + out[i] = __bfloat162float(bf[i]); + } +} + +// Store N floats as precision_t via 128-bit writes (float4 for f32, uint4 for bf16) +__device__ __forceinline__ void adv_vec_store(float* ptr, const float* vals) { + *reinterpret_cast(ptr) = make_float4(vals[0], vals[1], vals[2], vals[3]); +} + +__device__ __forceinline__ void adv_vec_store(__hip_bfloat16* ptr, const float* vals) { + // N=8 for bf16: all 8 elements fit in one uint4 (128 bits) + __hip_bfloat16 tmp[8]; + #pragma unroll + for (int i = 0; i < 8; i++) tmp[i] = __float2bfloat16(vals[i]); + *reinterpret_cast(ptr) = *reinterpret_cast(tmp); +} + +__device__ __forceinline__ void puff_advantage_row_vec( + const precision_t* values, const precision_t* rewards, const precision_t* dones, + const precision_t* importance, precision_t* advantages, float gamma, float lambda, + float rho_clip, float c_clip, int horizon) { + constexpr int N = 16 / sizeof(precision_t); + + float lastpufferlam = 0.0f; + int num_chunks = horizon / N; + + float next_value = to_float(values[horizon - 1]); + float next_done = to_float(dones[horizon - 1]); + float next_reward = to_float(rewards[horizon - 1]); + + for (int chunk = num_chunks - 1; chunk >= 0; chunk--) { + int base = chunk * N; + + float v[N], r[N], d[N], imp[N]; + adv_vec_load(values + base, v); + adv_vec_load(rewards + base, r); + adv_vec_load(dones + base, d); + adv_vec_load(importance + base, imp); + + float adv[N] = {0}; + int start_idx = (chunk == num_chunks - 1) ? (N - 2) : (N - 1); + + #pragma unroll + for (int i = start_idx; i >= 0; i--) { + float nextnonterminal = 1.0f - next_done; + float rho_t = fminf(imp[i], rho_clip); + float c_t = fminf(imp[i], c_clip); + float delta = rho_t * (next_reward + gamma * next_value * nextnonterminal - v[i]); + lastpufferlam = delta + gamma * lambda * c_t * lastpufferlam * nextnonterminal; + adv[i] = lastpufferlam; + next_value = v[i]; + next_done = d[i]; + next_reward = r[i]; + } + + adv_vec_store(advantages + base, adv); + } +} + +__global__ void puff_advantage(const precision_t* values, const precision_t* rewards, + const precision_t* dones, const precision_t* importance, precision_t* advantages, float gamma, + float lambda, float rho_clip, float c_clip, int num_steps, int horizon) { + int row = blockIdx.x*blockDim.x + threadIdx.x; + if (row >= num_steps) { + return; + } + int offset = row*horizon; + puff_advantage_row_vec(values + offset, rewards + offset, dones + offset, + importance + offset, advantages + offset, gamma, lambda, rho_clip, c_clip, horizon); +} + +__global__ void puff_advantage_scalar(const precision_t* values, const precision_t* rewards, + const precision_t* dones, const precision_t* importance, precision_t* advantages, float gamma, + float lambda, float rho_clip, float c_clip, int num_steps, int horizon) { + int row = blockIdx.x*blockDim.x + threadIdx.x; + if (row >= num_steps) { + return; + } + int offset = row*horizon; + puff_advantage_row_scalar(values + offset, rewards + offset, dones + offset, + importance + offset, advantages + offset, gamma, lambda, rho_clip, c_clip, horizon); +} + +void puff_advantage_cuda(PrecisionTensor& values, PrecisionTensor& rewards, + PrecisionTensor& dones, PrecisionTensor& importance, PrecisionTensor& advantages, + float gamma, float lambda, float rho_clip, float c_clip, hipStream_t stream) { + int num_steps = values.shape[0], horizon = values.shape[1]; + int blocks = grid_size(num_steps); + constexpr int N = 16 / sizeof(precision_t); + auto kernel = (horizon % N == 0) ? puff_advantage : puff_advantage_scalar; + kernel<<>>( + values.data, rewards.data, dones.data, importance.data, + advantages.data, gamma, lambda, rho_clip, c_clip, num_steps, horizon); +} + +// Zero advantages on frozen-bank rows so prio_replay never samples them. Frozen +// rollout rows hold actions/logprobs from the frozen policy — training the +// primary's PPO on them produces garbage ratios and poisoned gradients. +__global__ void zero_frozen_advantages_kernel(precision_t* advantages, + int agents_per_buffer, int primary_per_buffer, int total_rows, int horizon) { + int idx = blockIdx.x * blockDim.x + threadIdx.x; + int total = total_rows * horizon; + if (idx >= total) return; + int row = idx / horizon; + int rel = row % agents_per_buffer; + if (rel >= primary_per_buffer) { + advantages[idx] = from_float(0.0f); + } +} + +void zero_frozen_advantages_cuda(PrecisionTensor& advantages, + int agents_per_buffer, int primary_per_buffer, hipStream_t stream) { + int total_rows = advantages.shape[0]; + int horizon = advantages.shape[1]; + int total = total_rows * horizon; + zero_frozen_advantages_kernel<<>>( + advantages.data, agents_per_buffer, primary_per_buffer, total_rows, horizon); +} + +// Minor copy bandwidth optimizations +__global__ void index_copy(char* __restrict__ dst, const int* __restrict__ idx, + const char* __restrict__ src, int num_idx, int row_bytes) { + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < num_idx) { + int dst_row = idx[i]; + memcpy(dst + (int64_t)dst_row * row_bytes, src + (int64_t)i * row_bytes, row_bytes); + } +} + +__device__ __forceinline__ void copy_values_adv_returns( + const precision_t* __restrict__ src_values, precision_t* __restrict__ dst_values, + const precision_t* __restrict__ src_advantages, precision_t* __restrict__ dst_advantages, + precision_t* __restrict__ dst_returns, + int src_row, int dst_row, int horizon) { + int srh = (int64_t)src_row * horizon; + int drh = (int64_t)dst_row * horizon; + const precision_t* s_values = src_values + srh; + const precision_t* s_adv = src_advantages + srh; + precision_t* d_values = dst_values + drh; + precision_t* d_adv = dst_advantages + drh; + precision_t* d_returns = dst_returns + drh; + for (int i = threadIdx.x; i < horizon; i += blockDim.x) { + precision_t val = s_values[i]; + precision_t adv = s_adv[i]; + d_values[i] = val; + d_adv[i] = adv; + d_returns[i] = from_float(to_float(val) + to_float(adv)); + } +} + +__global__ void select_copy(RolloutBuf rollouts, TrainGraph graph, + const int* __restrict__ idx, const precision_t* __restrict__ advantages, + const float* __restrict__ mb_prio) { + int mb = blockIdx.x; + int ch = blockIdx.y; + int src_row = idx[mb]; + + // Compute row byte counts from tensor shapes + int obs_row_bytes = (numel(rollouts.observations.shape) / rollouts.observations.shape[0]) * sizeof(precision_t); + int act_row_bytes = (numel(rollouts.actions.shape) / rollouts.actions.shape[0]) * sizeof(precision_t); + int lp_row_bytes = (numel(rollouts.logprobs.shape) / rollouts.logprobs.shape[0]) * sizeof(precision_t); + int horizon = rollouts.values.shape[1]; + + switch (ch) { + case 0: + copy_bytes((const char*)rollouts.observations.data, (char*)graph.mb_obs.data, src_row, mb, obs_row_bytes); + break; + case 1: + copy_bytes((const char*)rollouts.actions.data, (char*)graph.mb_actions.data, src_row, mb, act_row_bytes); + break; + case 2: + copy_bytes((const char*)rollouts.logprobs.data, (char*)graph.mb_logprobs.data, src_row, mb, lp_row_bytes); + break; + case 3: + copy_values_adv_returns(rollouts.values.data, graph.mb_values.data, + advantages, graph.mb_advantages.data, + graph.mb_returns.data, src_row, mb, horizon); + break; + case 4: + if (threadIdx.x == 0) { + graph.mb_prio.data[mb] = from_float(mb_prio[mb]); + } + break; + case 5: + if (graph.mb_action_mask.data != nullptr) { + int mask_row_bytes = (numel(rollouts.action_mask.shape) + / rollouts.action_mask.shape[0]) * sizeof(precision_t); + copy_bytes((const char*)rollouts.action_mask.data, + (char*)graph.mb_action_mask.data, src_row, mb, mask_row_bytes); + } + break; + } +} + +inline float cosine_annealing(float lr_base, float lr_min, long t, long T) { + if (T == 0) return lr_base; + float ratio = (double )t / (double) T; + ratio = std::max(0.0f, std::min(1.0f, ratio)); + return lr_min + 0.5f*(lr_base - lr_min)*(1.0f + std::cos(M_PI * ratio)); +} + +void train_impl(PuffeRL& pufferl) { + // Update to HypersT& p + HypersT& hypers = pufferl.hypers; + + hipEventRecord(pufferl.profile.events[0]); // pre-loop start + hipStream_t train_stream = pufferl.default_stream; + + // Transpose from rollout layout (T, B, ...) to train layout (B, T, ...) + RolloutBuf& src = pufferl.rollouts; + RolloutBuf& rollouts = pufferl.train_rollouts; + PrecisionTensor& advantages_puf = pufferl.advantages_puf; + + int T = src.observations.shape[0], B = src.observations.shape[1]; + int obs_size = (ndim(src.observations.shape) >= 3) ? src.observations.shape[2] : 1; + int num_atns = (ndim(src.actions.shape) >= 3) ? src.actions.shape[2] : 1; + + transpose_102<<>>( + rollouts.observations.data, src.observations.data, T, B, obs_size); + transpose_102<<>>( + rollouts.actions.data, src.actions.data, T, B, num_atns); + transpose_102<<>>( + rollouts.logprobs.data, src.logprobs.data, T, B, 1); + transpose_102<<>>( + rollouts.rewards.data, src.rewards.data, T, B, 1); + transpose_102<<>>( + rollouts.terminals.data, src.terminals.data, T, B, 1); + transpose_102<<>>( + rollouts.ratio.data, src.ratio.data, T, B, 1); + transpose_102<<>>( + rollouts.values.data, src.values.data, T, B, 1); + if (src.action_mask.data != nullptr) { + int mask_size = src.action_mask.shape[2]; + transpose_102<<>>( + rollouts.action_mask.data, src.action_mask.data, T, B, mask_size); + } + + // We hard-clamp rewards to -1, 1. Our envs are mostly designed to respect this range + clamp_precision_kernel<<>>( + rollouts.rewards.data, -1.0f, 1.0f, numel(rollouts.rewards.shape)); + + // Set importance weights to 1.0 + fill_precision_kernel<<>>( + rollouts.ratio.data, from_float(1.0f), numel(rollouts.ratio.shape)); + + // Inline any of these only used once + int minibatch_size = hypers.minibatch_size; + int batch_size = hypers.total_agents * hypers.horizon; + int minibatch_segments = minibatch_size / hypers.horizon; + float prio_beta0 = hypers.prio_beta0; + float prio_alpha = hypers.prio_alpha; + bool anneal_lr = hypers.anneal_lr; + int current_epoch = pufferl.epoch; + + Muon* muon = &pufferl.muon; + int total_epochs = hypers.total_timesteps / batch_size; + if (anneal_lr) { + float lr_min = hypers.min_lr_ratio * hypers.lr; + float lr = cosine_annealing(hypers.lr, lr_min, current_epoch, total_epochs); + hipMemcpy(muon->lr_ptr, &lr, sizeof(float), hipMemcpyHostToDevice); + } + + // Annealed entropy coefficient — same cosine shape as lr. With PG signal + // alive, the entropy bonus that kept early-training exploratory becomes + // load-bearing dead weight late in training; cosine-decay frees the policy + // to commit harder on what it has already learned. + float current_ent_coef = hypers.ent_coef; + if (hypers.anneal_ent_coef) { + float ent_min = hypers.min_ent_coef_ratio * hypers.ent_coef; + current_ent_coef = cosine_annealing(hypers.ent_coef, ent_min, + current_epoch, total_epochs); + } + // copy ent_coef to the device buffer read by the loss kernel + hipMemcpy(pufferl.ppo_bufs_puf.ent_coef.data, ¤t_ent_coef, + sizeof(float), hipMemcpyHostToDevice); + + // Annealed priority exponent + float anneal_beta = prio_beta0 + (1.0f - prio_beta0) * prio_alpha * (float)current_epoch/(float)total_epochs; + TrainGraph& graph = pufferl.train_buf; + hipEventRecord(pufferl.profile.events[1]); // pre-loop end + + int total_minibatches = hypers.replay_ratio * batch_size / hypers.minibatch_size; + for (int mb = 0; mb < total_minibatches; ++mb) { + hipEventRecord(pufferl.profile.events[2]); // start of misc (overwritten each iter) + puf_zero(&advantages_puf, train_stream); + + profile_begin("compute_advantage", hypers.profile); + puff_advantage_cuda(rollouts.values, rollouts.rewards, rollouts.terminals, + rollouts.ratio, advantages_puf, hypers.gamma, hypers.gae_lambda, + hypers.vtrace_rho_clip, hypers.vtrace_c_clip, train_stream); + if (pufferl.num_frozen_banks > 0 && pufferl.bank_layout != NULL) { + int apb = hypers.total_agents / hypers.num_buffers; + zero_frozen_advantages_cuda(advantages_puf, apb, + pufferl.bank_layout[1], train_stream); + } + profile_end(hypers.profile); + + profile_begin("compute_prio", hypers.profile); + // Use the training RNG offset slot (last slot, index num_buffers) + long* train_rng_offset = pufferl.rng_offset_puf.data + hypers.num_buffers; + prio_replay_cuda(advantages_puf, prio_alpha, minibatch_segments, + hypers.total_agents, anneal_beta, + pufferl.prio_bufs, pufferl.seed, train_rng_offset, train_stream); + profile_end(hypers.profile); + + profile_begin("train_select_and_copy", hypers.profile); + if (hypers.reset_state) puf_zero(&graph.mb_state, train_stream); + { + RolloutBuf sel_src = rollouts; + sel_src.values = rollouts.values; + int mb_segs = pufferl.prio_bufs.idx.shape[0]; + int channels = (graph.mb_action_mask.data != nullptr) ? 6 : 5; + select_copy<<>>( + sel_src, graph, pufferl.prio_bufs.idx.data, + advantages_puf.data, pufferl.prio_bufs.mb_prio.data); + } + profile_end(hypers.profile); + + hipEventRecord(pufferl.profile.events[3]); // end misc / start forward + profile_begin("train_forward_backward", hypers.profile); + if (pufferl.train_captured) { + hipGraphLaunch(pufferl.train_cudagraph, train_stream); + } else { + bool capturing = pufferl.train_warmup == hypers.cudagraphs; + if (capturing) { + assert(hipStreamBeginCapture(train_stream, hipStreamCaptureModeGlobal) == hipSuccess + && "hipStreamBeginCapture failed"); + } + + hipStream_t stream = train_stream; + PrecisionTensor obs_puf = graph.mb_obs; + PrecisionTensor state_puf = graph.mb_state; + PrecisionTensor dec_puf = policy_forward_train(&pufferl.policy, pufferl.weights, pufferl.train_activations, obs_puf, state_puf, stream); + DecoderWeights* dw_train = (DecoderWeights*)pufferl.weights.decoder; + PrecisionTensor p_logstd; + if (dw_train->continuous) { + p_logstd = dw_train->logstd; + } + + ppo_loss_fwd_bwd(dec_puf, p_logstd, graph, + pufferl.act_sizes_puf, pufferl.losses_puf, + hypers.clip_coef, hypers.vf_clip_coef, hypers.vf_coef, + pufferl.ppo_bufs_puf.ent_coef.data, + pufferl.ppo_bufs_puf, pufferl.is_continuous, stream); + + FloatTensor grad_logits_puf = pufferl.ppo_bufs_puf.grad_logits; + FloatTensor grad_logstd_puf = pufferl.is_continuous ? pufferl.ppo_bufs_puf.grad_logstd : FloatTensor(); + FloatTensor grad_values_puf = pufferl.ppo_bufs_puf.grad_values; + policy_backward(&pufferl.policy, pufferl.weights, pufferl.train_activations, + grad_logits_puf, grad_logstd_puf, grad_values_puf, stream); + + muon_step(&pufferl.muon, pufferl.master_weights, pufferl.grad_puf, hypers.max_grad_norm, stream); + if (USE_BF16) { + int n = numel(pufferl.param_puf.shape); + cast<<>>( + pufferl.param_puf.data, pufferl.master_weights.data, n); + } + if (capturing) { + hipGraph_t _graph; + assert(hipStreamEndCapture(train_stream, &_graph) == hipSuccess + && "hipStreamEndCapture failed"); + assert(hipGraphInstantiateWithFlags(&pufferl.train_cudagraph, _graph, 0) == hipSuccess + && "hipGraphInstantiate failed"); + assert(hipGraphDestroy(_graph) == hipSuccess && "hipGraphDestroy failed"); + hipDeviceSynchronize(); + pufferl.train_captured = true; + } + pufferl.train_warmup++; + } + profile_end(hypers.profile); + + // This version is consistent with PufferLib 3.0. One of the major algorithmic + // questions remaining is how and when to update value and advantage estimates. + { + int num_idx = numel(pufferl.prio_bufs.idx.shape); + int row_bytes = (numel(graph.mb_ratio.shape) / graph.mb_ratio.shape[0]) * sizeof(precision_t); + index_copy<<>>( + (char*)rollouts.ratio.data, pufferl.prio_bufs.idx.data, + (const char*)graph.mb_ratio.data, num_idx, row_bytes); + } + { + int num_idx = numel(pufferl.prio_bufs.idx.shape); + int row_bytes = graph.mb_newvalue.shape[1] * sizeof(precision_t); + index_copy<<>>( + (char*)rollouts.values.data, pufferl.prio_bufs.idx.data, + (const char*)graph.mb_newvalue.data, num_idx, row_bytes); + } + hipEventRecord(pufferl.profile.events[4]); // end forward + } + pufferl.epoch += 1; + + hipStreamSynchronize(pufferl.default_stream); + + if (total_minibatches > 0) { + float ms; + // Pre-loop setup (transpose, advantage, allocs) + hipEventElapsedTime(&ms, pufferl.profile.events[0], pufferl.profile.events[1]); + pufferl.profile.accum[PROF_TRAIN_MISC] += ms; + // In-loop misc (last iteration, representative) scaled by count + hipEventElapsedTime(&ms, pufferl.profile.events[2], pufferl.profile.events[3]); + pufferl.profile.accum[PROF_TRAIN_MISC] += ms * total_minibatches; + // In-loop forward (last iteration, representative) scaled by count + hipEventElapsedTime(&ms, pufferl.profile.events[3], pufferl.profile.events[4]); + pufferl.profile.accum[PROF_TRAIN_FORWARD] += ms * total_minibatches; + } + +} + +// Build a Policy value for a given env + arch. Encoder/decoder algorithms are +// fixed by the env; hidden_size/num_layers/horizon parameterize shape. Policy +// has no heap state so this returns by value; callers store it wherever. +static Policy build_policy(const char* env_name, int input_size, int hidden_size, + int num_layers, int decoder_output_size, int act_n, + bool is_continuous, int horizon) { + Encoder encoder = { + .forward = encoder_forward, + .backward = encoder_backward, + .init_weights = encoder_init_weights, + .reg_params = encoder_reg_params, + .reg_train = encoder_reg_train, + .reg_rollout = encoder_reg_rollout, + .create_weights = encoder_create_weights, + .free_weights = encoder_free_weights, + .free_activations = encoder_free_activations, + .in_dim = input_size, .out_dim = hidden_size, + .activation_size = sizeof(EncoderActivations), + }; + create_custom_encoder(env_name, &encoder); + Decoder decoder = { + .forward = decoder_forward, + .backward = decoder_backward, + .init_weights = decoder_init_weights, + .reg_params = decoder_reg_params, + .reg_train = decoder_reg_train, + .reg_rollout = decoder_reg_rollout, + .create_weights = decoder_create_weights, + .free_weights = decoder_free_weights, + .free_activations = decoder_free_activations, + .hidden_dim = hidden_size, .output_dim = decoder_output_size, .continuous = is_continuous, + .activation_size = (int)sizeof(DecoderActivations), + }; + create_custom_decoder(env_name, &decoder); + Network network = { + .forward = mingru_forward, + .forward_train = mingru_forward_train, + .backward = mingru_backward, + .init_weights = mingru_init_weights, + .reg_params = mingru_reg_params, + .reg_train = mingru_reg_train, + .reg_rollout = mingru_reg_rollout, + .create_weights = mingru_create_weights, + .free_weights = mingru_free_weights, + .free_activations = mingru_free_activations, + .hidden = hidden_size, .num_layers = num_layers, .horizon = horizon, + }; + return Policy{ + .encoder = encoder, .decoder = decoder, .network = network, + .input_dim = input_size, .hidden_dim = hidden_size, .output_dim = decoder_output_size, + .num_atns = act_n, + }; +} + +// Allocate a fresh frozen WeightBank with its own Policy (may differ in +// hidden_size/num_layers from primary). slice_size = how many agents per buffer +// this bank will own. Weights are uninitialized — caller must load before use. +static void weight_bank_create_for_pufferl(WeightBank* bank, PuffeRL* pufferl, + int slice_size, int hidden_size, int num_layers) { + int num_buffers = pufferl->hypers.num_buffers; + + // Rebuild arch-varying Policy from env metadata already on pufferl. + int input_size = pufferl->env.obs.shape[1]; + int num_action_heads = pufferl->env.actions.shape[1]; + int* raw_act_sizes = get_act_sizes(); + int act_n = 0; + for (int i = 0; i < num_action_heads; i++) act_n += raw_act_sizes[i]; + int decoder_output_size = pufferl->is_continuous ? num_action_heads : act_n; + bank->policy = build_policy(pufferl->env_name.c_str(), input_size, hidden_size, + num_layers, decoder_output_size, act_n, pufferl->is_continuous, pufferl->hypers.horizon); + bank->hidden_size = hidden_size; + bank->num_layers = num_layers; + + Allocator* params = &bank->params_alloc; + Allocator* acts = &bank->acts_alloc; + + bank->slice_size = slice_size; + bank->weights = policy_weights_create(&bank->policy, params); + bank->buffer_activations = (PolicyActivations*)calloc(num_buffers, sizeof(PolicyActivations)); + bank->buffer_states = (PrecisionTensor*)calloc(num_buffers, sizeof(PrecisionTensor)); + for (int i = 0; i < num_buffers; i++) { + bank->buffer_activations[i] = policy_reg_rollout(&bank->policy, bank->weights, acts, slice_size); + bank->buffer_states[i] = {.shape = {num_layers, slice_size, hidden_size}}; + alloc_register(acts, &bank->buffer_states[i]); + } + + alloc_create(params); + alloc_create(acts); + + bank->param_puf = {.data = (precision_t*)params->mem, .shape = {params->total_elems}}; + if (USE_BF16) { + bank->master_weights = {.shape = {params->total_elems}}; + hipMalloc(&bank->master_weights.data, params->total_elems * sizeof(float)); + } else { + bank->master_weights = {.data = (float*)bank->param_puf.data, .shape = {params->total_elems}}; + } +} + +// Mirror of weight_bank_create_for_pufferl. Frees the bank's weights, per-buffer +// activations, allocators, and master_weights (BF16 only). Does not free the +// WeightBank struct itself — caller owns that. +static void weight_bank_destroy(WeightBank* bank, PuffeRL* pufferl) { + int num_buffers = pufferl->hypers.num_buffers; + policy_weights_free(&bank->policy, &bank->weights); + if (bank->buffer_activations != NULL) { + for (int i = 0; i < num_buffers; i++) { + policy_activations_free(&bank->policy, bank->buffer_activations[i]); + } + free(bank->buffer_activations); + } + free(bank->buffer_states); + alloc_free(&bank->params_alloc); + alloc_free(&bank->acts_alloc); + if (USE_BF16 && bank->master_weights.data != NULL) { + hipFree(bank->master_weights.data); + } +} + +// Append a fresh frozen bank with the given per-buffer slice size; returns its +// index. Rebuilds bank_layout sequentially (primary first, then frozen banks in +// add order). Must be called BEFORE cudagraph capture (pointers get baked in). +extern "C" int pufferl_add_frozen_bank(PuffeRL* pufferl, int slice_size, + int hidden_size, int num_layers) { + int idx = pufferl->num_frozen_banks; + pufferl->frozen_banks = (WeightBank*)realloc( + pufferl->frozen_banks, (idx + 1) * sizeof(WeightBank)); + memset(&pufferl->frozen_banks[idx], 0, sizeof(WeightBank)); + weight_bank_create_for_pufferl(&pufferl->frozen_banks[idx], pufferl, + slice_size, hidden_size, num_layers); + pufferl->num_frozen_banks++; + + // Rebuild sequential layout from declared slice_sizes. + int agents_per_buffer = pufferl->vec->total_agents / pufferl->hypers.num_buffers; + int frozen_total = 0; + for (int b = 0; b < pufferl->num_frozen_banks; b++) { + frozen_total += pufferl->frozen_banks[b].slice_size; + } + if (frozen_total > agents_per_buffer) { + fprintf(stderr, "pufferl_add_frozen_bank: total frozen slice (%d) exceeds " + "agents_per_buffer (%d)\n", frozen_total, agents_per_buffer); + } + int num_banks = 1 + pufferl->num_frozen_banks; + pufferl->bank_layout = (int*)realloc(pufferl->bank_layout, (num_banks + 1) * sizeof(int)); + pufferl->bank_layout[0] = 0; + pufferl->bank_layout[1] = agents_per_buffer - frozen_total; // primary + int cumul = pufferl->bank_layout[1]; + for (int b = 0; b < pufferl->num_frozen_banks; b++) { + cumul += pufferl->frozen_banks[b].slice_size; + pufferl->bank_layout[2 + b] = cumul; + } + return idx; +} + +// Load a frozen bank's weights from a file (same format as save_weights — flat fp32). +// Safe to call between rollouts (in-place hipMemcpy; cudagraphs hold the pointer, +// not a copy of the data). +extern "C" void pufferl_load_frozen_bank(PuffeRL* pufferl, int bank_idx, const char* path) { + if (bank_idx < 0 || bank_idx >= pufferl->num_frozen_banks) { + fprintf(stderr, "pufferl_load_frozen_bank: bank_idx %d out of range\n", bank_idx); + return; + } + WeightBank* bank = &pufferl->frozen_banks[bank_idx]; + int64_t nbytes = numel(bank->master_weights.shape) * sizeof(float); + FILE* f = fopen(path, "rb"); + if (!f) { + fprintf(stderr, "pufferl_load_frozen_bank: failed to open %s\n", path); + return; + } + fseek(f, 0, SEEK_END); + long file_size = ftell(f); + fseek(f, 0, SEEK_SET); + if (file_size != nbytes) { + fprintf(stderr, "pufferl_load_frozen_bank: size mismatch (expected %lld, got %ld)\n", + (long long)nbytes, file_size); + fclose(f); + return; + } + std::vector buf(nbytes); + size_t nread = fread(buf.data(), 1, nbytes, f); + fclose(f); + if ((int64_t)nread != nbytes) { + fprintf(stderr, "pufferl_load_frozen_bank: short read on %s\n", path); + return; + } + hipMemcpy(bank->master_weights.data, buf.data(), nbytes, hipMemcpyHostToDevice); + if (USE_BF16) { + int n = numel(bank->param_puf.shape); + cast<<default_stream>>>( + bank->param_puf.data, bank->master_weights.data, n); + } + hipDeviceSynchronize(); +} + +// Set the agent permutation. Validates that the perm respects buffer boundaries: +// each buffer's range [buf_start, buf_start+buf_size) must map onto itself (no +// cross-buffer writes, since each worker only owns its physical chunk). +extern "C" void pufferl_set_agent_perm(PuffeRL* pufferl, const int* perm) { + int total = pufferl->vec->total_agents; + int num_buffers = pufferl->hypers.num_buffers; + int buf_size = total / num_buffers; + for (int b = 0; b < num_buffers; b++) { + int lo = b * buf_size; + int hi = lo + buf_size; + for (int i = lo; i < hi; i++) { + if (perm[i] < lo || perm[i] >= hi) { + fprintf(stderr, + "pufferl_set_agent_perm: perm[%d]=%d crosses buffer %d range [%d,%d)\n", + i, perm[i], b, lo, hi); + return; + } + } + } + static_vec_set_perm(pufferl->vec, perm); +} + +// Set per-env tags (e.g. selfplay vs historical). tags array length must equal +// pufferl_num_envs(). Also clears each env's boundary_reached flag. +extern "C" void pufferl_set_env_tags(PuffeRL* pufferl, const int* tags) { + static_vec_set_env_tags(pufferl->vec, tags); +} + +// Returns count of envs with tag == tag_value AND boundary_reached. If +// reset_flags != 0, clears boundary_reached only on envs whose tag matches +// tag_value (so multi-bank swaps don't trample each other's alignment). +extern "C" int pufferl_count_aligned(PuffeRL* pufferl, int tag_value, int reset_flags) { + return static_vec_count_aligned(pufferl->vec, tag_value, reset_flags); +} + +extern "C" int pufferl_num_envs(PuffeRL* pufferl) { + return pufferl->vec->size; +} + +std::unique_ptr create_pufferl_impl(HypersT& hypers, + const std::string& env_name, Dict* vec_kwargs, Dict* env_kwargs) { + auto pufferl = std::make_unique(); + pufferl->hypers = hypers; + pufferl->nccl_comm = nullptr; + pufferl->default_stream = 0; + pufferl->env_name = env_name; + + hipSetDevice(hypers.gpu_id); + + // Multi-GPU: initialize NCCL + if (hypers.world_size > 1) { + if (hypers.nccl_id.size() != sizeof(ncclUniqueId)) + throw std::runtime_error("nccl_id must be " + std::to_string(sizeof(ncclUniqueId)) + " bytes"); + ncclUniqueId nccl_id; + memcpy(&nccl_id, hypers.nccl_id.data(), sizeof(nccl_id)); + ncclCommInitRank(&pufferl->nccl_comm, hypers.world_size, nccl_id, hypers.rank); + printf("Rank %d/%d: NCCL initialized\n", hypers.rank, hypers.world_size); + } + + ulong seed = hypers.seed + hypers.rank; + pufferl->seed = seed; + + // Load environment first to get input_size and action info from env + // Create environments and set up action sizes + StaticVec* vec = create_environments(hypers.num_buffers, hypers.total_agents, + env_name, vec_kwargs, env_kwargs, pufferl->env); + pufferl->vec = vec; + + // Sanity check action space + int num_action_heads = pufferl->env.actions.shape[1]; + int* raw_act_sizes = get_act_sizes(); // CPU int32 pointer from env + int act_n = 0; + int num_continuous = 0; + int num_discrete = 0; + for (int i = 0; i < num_action_heads; i++) { + int val = raw_act_sizes[i]; + if (val == 1) { + num_continuous++; + } else { + num_discrete++; + } + act_n += val; + } + assert((num_continuous == 0 || num_discrete == 0) && + "Mixed continuous/discrete action spaces not supported"); + pufferl->is_continuous = (num_continuous > 0); + if (pufferl->is_continuous) { + printf("Detected continuous action space with %d dimensions\n", num_action_heads); + } else { + printf("Detected discrete action space with %d heads\n", num_action_heads); + } + + // Create profiling events + for (int i = 0; i < NUM_TRAIN_EVENTS; i++) { + hipEventCreate(&pufferl->profile.events[i]); + } + memset(pufferl->profile.accum, 0, sizeof(pufferl->profile.accum)); + nvmlInit(); + nvmlDeviceGetHandleByIndex(hypers.gpu_id, &pufferl->nvml_device); + + // Create policy + int input_size = pufferl->env.obs.shape[1]; + int hidden_size = hypers.hidden_size; + int num_layers = hypers.num_layers; + bool is_continuous = pufferl->is_continuous; + int decoder_output_size = is_continuous ? num_action_heads : act_n; + int minibatch_segments = hypers.minibatch_size / hypers.horizon; + int inf_batch = vec->total_agents / hypers.num_buffers; + int B_TT = minibatch_segments * hypers.horizon; + int horizon = hypers.horizon; + int total_agents = vec->total_agents; + int batch = total_agents / hypers.num_buffers; + int num_buffers = hypers.num_buffers; + + pufferl->policy = build_policy(env_name.c_str(), input_size, hidden_size, + num_layers, decoder_output_size, act_n, is_continuous, hypers.horizon); + + // Create and allocate params + Allocator* params = &pufferl->params_alloc; + Allocator* acts = &pufferl->activations_alloc; + Allocator* grads = &pufferl->grads_alloc; + + // Buffers for weights, grads, and activations + pufferl->weights = policy_weights_create(&pufferl->policy, params); + pufferl->train_activations = policy_reg_train(&pufferl->policy, pufferl->weights, acts, grads, B_TT); + pufferl->buffer_activations = (PolicyActivations*)calloc(num_buffers, sizeof(PolicyActivations)); + pufferl->buffer_states = (PrecisionTensor*)calloc(num_buffers, sizeof(PrecisionTensor)); + for (int i = 0; i < num_buffers; i++) { + pufferl->buffer_activations[i] = policy_reg_rollout( + &pufferl->policy, pufferl->weights, acts, inf_batch); + pufferl->buffer_states[i] = { + .shape = {num_layers, batch, hidden_size}, + }; + alloc_register(acts, &pufferl->buffer_states[i]); + } + int mask_size = pufferl->vec->action_mask_size; + register_rollout_buffers(pufferl->rollouts, + acts, horizon, total_agents, input_size, num_action_heads, mask_size); + register_train_buffers(pufferl->train_buf, + acts, minibatch_segments, horizon, input_size, + hidden_size, num_action_heads, num_layers, mask_size); + register_rollout_buffers(pufferl->train_rollouts, + acts, total_agents, horizon, input_size, num_action_heads, mask_size); + register_ppo_buffers(pufferl->ppo_bufs_puf, + acts, minibatch_segments, hypers.horizon, decoder_output_size, is_continuous); + register_prio_buffers(pufferl->prio_bufs, + acts, hypers.total_agents, minibatch_segments); + + // Extra cuda buffers just reuse activ allocator + pufferl->rng_offset_puf = {.shape = {num_buffers + 1}}; + alloc_register(acts, &pufferl->rng_offset_puf); + + pufferl->act_sizes_puf = {.shape = {num_action_heads}}; + alloc_register(acts, &pufferl->act_sizes_puf); + + pufferl->losses_puf = {.shape = {NUM_LOSSES}}; + alloc_register(acts, &pufferl->losses_puf); + + pufferl->advantages_puf = {.shape = {total_agents, horizon}}; + alloc_register(acts, &pufferl->advantages_puf); + + muon_init(&pufferl->muon, params, hypers.lr, hypers.beta1, hypers.eps, 0.0, acts); + pufferl->muon.nccl_comm = pufferl->nccl_comm; + pufferl->muon.world_size = hypers.world_size; + + // All buffers allocated here + if (alloc_create(params) != hipSuccess) { + return nullptr; + } + if (alloc_create(grads) != hipSuccess) { + return nullptr; + } + if (alloc_create(acts) != hipSuccess) { + return nullptr; + } + + pufferl->grad_puf = {.data = (precision_t*)grads->mem, .shape = {grads->total_elems}}; + pufferl->param_puf = {.data = (precision_t*)params->mem, .shape = {params->total_elems}}; + + ulong init_seed = hypers.seed; + policy_init_weights(&pufferl->policy, pufferl->weights, &init_seed, pufferl->default_stream); + pufferl->master_weights = {.data = (float*)pufferl->param_puf.data, .shape = {params->total_elems}}; + if (USE_BF16) { + pufferl->master_weights = {.shape = {params->total_elems}}; + hipMalloc(&pufferl->master_weights.data, params->total_elems * sizeof(float)); + int n = numel(pufferl->param_puf.shape); + cast<<default_stream>>>( + pufferl->master_weights.data, pufferl->param_puf.data, n); + } + + // Per-buffer persistent RNG states + int agents_per_buf = total_agents / num_buffers; + pufferl->rng_states = (hiprandStatePhilox4_32_10_t**)calloc(num_buffers, sizeof(hiprandStatePhilox4_32_10_t*)); + for (int i = 0; i < num_buffers; i++) { + hipMalloc(&pufferl->rng_states[i], agents_per_buf * sizeof(hiprandStatePhilox4_32_10_t)); + rng_init<<>>( + pufferl->rng_states[i], pufferl->seed + i, agents_per_buf); + } + + // Post-create initialization + hipMemcpy(pufferl->act_sizes_puf.data, raw_act_sizes, num_action_heads * sizeof(int), hipMemcpyHostToDevice); + hipMemset(pufferl->losses_puf.data, 0, NUM_LOSSES * sizeof(float)); + float one = 1.0f; + hipMemcpy(pufferl->ppo_bufs_puf.grad_loss.data, &one, sizeof(float), hipMemcpyHostToDevice); + muon_post_create(&pufferl->muon); + + // Set up frozen banks declared in vec_kwargs (num_frozen_banks + + // frozen_bank_pct: each bank gets floor(agents_per_buffer * pct) agents). + // Must happen BEFORE cudagraph capture so the graph bakes in their pointers + // and per-bank loop iterations. + DictItem* nb_item = dict_get_unsafe(vec_kwargs, "num_frozen_banks"); + DictItem* fbp_item = dict_get_unsafe(vec_kwargs, "frozen_bank_pct"); + DictItem* fbh_item = dict_get_unsafe(vec_kwargs, "frozen_bank_hidden_size"); + DictItem* fbl_item = dict_get_unsafe(vec_kwargs, "frozen_bank_num_layers"); + int num_frozen = nb_item ? (int)nb_item->value : 0; + float frozen_pct = fbp_item ? (float)fbp_item->value : 0.0f; + int frozen_hidden = fbh_item ? (int)fbh_item->value : hidden_size; + int frozen_layers = fbl_item ? (int)fbl_item->value : num_layers; + if (num_frozen > 0) { + int agents_per_buffer = total_agents / num_buffers; + int frozen_size = (int)((float)agents_per_buffer * frozen_pct); // truncates = floor for positive + int frozen_total = num_frozen * frozen_size; + if (frozen_size <= 0 || frozen_total > agents_per_buffer) { + fprintf(stderr, "create_pufferl: invalid frozen bank config " + "(num=%d, pct=%.4f -> size=%d, total=%d, agents_per_buffer=%d)\n", + num_frozen, frozen_pct, frozen_size, frozen_total, agents_per_buffer); + return nullptr; + } + // add_frozen_bank auto-builds the sequential bank_layout. + for (int b = 0; b < num_frozen; b++) { + pufferl_add_frozen_bank(pufferl.get(), frozen_size, frozen_hidden, frozen_layers); + } + } + + // Cudagraph rolluts and entire training step + if (hypers.cudagraphs >= 0) { + pufferl->fused_rollout_cudagraphs = (hipGraphExec_t*)calloc(horizon*num_buffers, sizeof(hipGraphExec_t)); + pufferl->train_warmup = 0; + + // Snapshot weights + optimizer state before init-time capture + long wb_bytes = numel(pufferl->master_weights.shape) * sizeof(float); + void* saved_weights; + hipMalloc(&saved_weights, wb_bytes); + hipMemcpy(saved_weights, pufferl->master_weights.data, wb_bytes, hipMemcpyDeviceToDevice); + void* saved_momentum; + hipMalloc(&saved_momentum, wb_bytes); + hipMemcpy(saved_momentum, pufferl->muon.mb_puf.data, wb_bytes, hipMemcpyDeviceToDevice); + + // Create per-buffer streams before capture so graphs are + // captured and replayed on the same streams. + pufferl->streams = (hipStream_t*)calloc(num_buffers, sizeof(hipStream_t)); + for (int i = 0; i < num_buffers; i++) { + hipStreamCreate(&pufferl->streams[i]); + vec->streams[i] = pufferl->streams[i]; + } + + hipStream_t saved_default = pufferl->default_stream; + hipStream_t saved_tl = tl_stream; + hipStream_t warmup_stream; + hipStreamCreate(&warmup_stream); + pufferl->default_stream = warmup_stream; + + for (pufferl->epoch = 0; pufferl->epoch <= hypers.cudagraphs; pufferl->epoch++) { + for (int i = 0; i < num_buffers * horizon; ++i) { + int buf = i % num_buffers; + tl_stream = pufferl->streams[buf]; + net_callback_wrapper(pufferl.get(), buf, i / num_buffers); + hipDeviceSynchronize(); + } + } + pufferl->rollout_captured = true; + + tl_stream = warmup_stream; + for (int i = 0; i <= hypers.cudagraphs; i++) { + train_impl(*pufferl); + } + + hipStreamSynchronize(warmup_stream); + hipDeviceSynchronize(); + pufferl->default_stream = saved_default; + tl_stream = saved_tl; + hipStreamDestroy(warmup_stream); + + // Restore weights + optimizer state corrupted by warmup/capture + hipMemcpy(pufferl->master_weights.data, saved_weights, wb_bytes, hipMemcpyDeviceToDevice); + hipFree(saved_weights); + hipMemcpy(pufferl->muon.mb_puf.data, saved_momentum, wb_bytes, hipMemcpyDeviceToDevice); + hipFree(saved_momentum); + if (USE_BF16) { + int n = numel(pufferl->param_puf.shape); + cast<<default_stream>>>( + pufferl->param_puf.data, pufferl->master_weights.data, n); + } + + // Re-init RNG states corrupted by warmup + for (int i = 0; i < num_buffers; i++) { + rng_init<<>>( + pufferl->rng_states[i], pufferl->seed + i, agents_per_buf); + } + hipDeviceSynchronize(); + + pufferl->epoch = 0; + pufferl->global_step = 0; + } + + // Create per-buffer streams if not already created by cudagraph path + if (!pufferl->streams) { + pufferl->streams = (hipStream_t*)calloc(num_buffers, sizeof(hipStream_t)); + for (int i = 0; i < num_buffers; i++) { + hipStreamCreate(&pufferl->streams[i]); + vec->streams[i] = pufferl->streams[i]; + } + } + + create_static_threads(vec, hypers.num_threads, horizon, pufferl.get(), + net_callback_wrapper, thread_init_wrapper); + static_vec_reset(vec); + + if (hypers.profile) { + hipDeviceSynchronize(); + hipProfilerStart(); + } + + double now = wall_clock(); + pufferl->start_time = now; + pufferl->last_log_time = now; + pufferl->last_log_step = 0; + + return pufferl; +} + +void close_impl(PuffeRL& pufferl) { + hipDeviceSynchronize(); + if (pufferl.hypers.profile) { + hipProfilerStop(); + } + + hipGraphExecDestroy(pufferl.train_cudagraph); + for (int i = 0; i < pufferl.hypers.horizon * pufferl.hypers.num_buffers; i++) { + hipGraphExecDestroy(pufferl.fused_rollout_cudagraphs[i]); + } + + policy_weights_free(&pufferl.policy, &pufferl.weights); + policy_activations_free(&pufferl.policy, pufferl.train_activations); + for (int buf = 0; buf < pufferl.hypers.num_buffers; buf++) { + policy_activations_free(&pufferl.policy, pufferl.buffer_activations[buf]); + } + + for (int i = 0; i < pufferl.hypers.num_buffers; i++) { + hipFree(pufferl.rng_states[i]); + } + free(pufferl.rng_states); + + if (USE_BF16) { + hipFree(pufferl.master_weights.data); + } + + alloc_free(&pufferl.params_alloc); + alloc_free(&pufferl.grads_alloc); + alloc_free(&pufferl.activations_alloc); + + for (int i = 0; i < pufferl.hypers.num_buffers; i++) { + hipStreamDestroy(pufferl.streams[i]); + } + for (int i = 0; i < NUM_TRAIN_EVENTS; i++) { + hipEventDestroy(pufferl.profile.events[i]); + } + nvmlShutdown(); + + static_vec_close(pufferl.vec); + + free(pufferl.buffer_states); + free(pufferl.buffer_activations); + free(pufferl.fused_rollout_cudagraphs); + free(pufferl.streams); + + for (int b = 0; b < pufferl.num_frozen_banks; b++) { + weight_bank_destroy(&pufferl.frozen_banks[b], &pufferl); + } + free(pufferl.frozen_banks); + free(pufferl.bank_layout); + + if (pufferl.nccl_comm != nullptr) { + ncclCommDestroy(pufferl.nccl_comm); + } +} diff --git a/src-hip/stub/nvml.h b/src-hip/stub/nvml.h new file mode 100644 index 0000000000..4f09a532e5 --- /dev/null +++ b/src-hip/stub/nvml.h @@ -0,0 +1,20 @@ +// Minimal NVML stub for AMD builds - utilization stats read as zeros. +#ifndef STUB_NVML_H +#define STUB_NVML_H +#include +typedef struct nvmlDevice_st* nvmlDevice_t; +typedef struct { unsigned int gpu; unsigned int memory; } nvmlUtilization_t; +typedef struct { unsigned long long total; unsigned long long free; unsigned long long used; } nvmlMemory_t; +#ifndef STUB_NVML_FREE_MEMBER +#define STUB_NVML_FREE_MEMBER +#endif +static inline int nvmlInit_v2(void) { return 0; } +static inline int nvmlInit(void) { return 0; } +static inline int nvmlShutdown(void) { return 0; } +static inline int nvmlDeviceGetHandleByIndex(unsigned int i, nvmlDevice_t* d) { (void)i; *d = 0; return 0; } +static inline int nvmlDeviceGetUtilizationRates(nvmlDevice_t d, nvmlUtilization_t* u) { (void)d; u->gpu = 0; u->memory = 0; return 0; } +static inline int nvmlDeviceGetMemoryInfo(nvmlDevice_t d, nvmlMemory_t* m) { (void)d; m->total = 0; m->used = 0; m->free = 0; return 0; } +#endif +// NVTX no-ops for AMD builds +static inline int nvtxRangePushA(const char* name) { (void)name; return 0; } +static inline void nvtxRangePop(void) {} diff --git a/src/vecenv.h b/src/vecenv.h index 42958d321e..a1c1dbe4b2 100644 --- a/src/vecenv.h +++ b/src/vecenv.h @@ -64,7 +64,12 @@ static inline void dict_set(Dict* dict, const char* key, double value) { } // Forward declare CUDA stream type +#ifdef __HIP_PLATFORM_AMD__ +#include +#define cudaStream_t hipStream_t +#else typedef struct CUstream_st* cudaStream_t; +#endif // Threading state typedef struct StaticThreading StaticThreading; @@ -188,6 +193,23 @@ typedef int cudaMemcpyKind; #define cudaHostAllocPortable 1 #define cudaStreamNonBlocking 1 +#ifdef __HIP_PLATFORM_AMD__ +#include +#define cudaMemcpyKind hipMemcpyKind +#define cudaHostAlloc hipHostAlloc +#define cudaMalloc hipMalloc +#define cudaMemcpy hipMemcpy +#define cudaMemcpyAsync hipMemcpyAsync +#define cudaMemset hipMemset +#define cudaFree hipFree +#define cudaFreeHost hipHostFree +#define cudaSetDevice hipSetDevice +#define cudaDeviceSynchronize hipDeviceSynchronize +#define cudaStreamSynchronize hipStreamSynchronize +#define cudaStreamCreateWithFlags hipStreamCreateWithFlags +#define cudaStreamQuery hipStreamQuery +#define cudaGetErrorString hipGetErrorString +#else extern cudaError_t cudaHostAlloc(void**, size_t, unsigned int); extern cudaError_t cudaMalloc(void**, size_t); extern cudaError_t cudaMemcpy(void*, const void*, size_t, cudaMemcpyKind); @@ -201,6 +223,7 @@ extern cudaError_t cudaStreamSynchronize(cudaStream_t); extern cudaError_t cudaStreamCreateWithFlags(cudaStream_t*, unsigned int); extern cudaError_t cudaStreamQuery(cudaStream_t); extern const char* cudaGetErrorString(cudaError_t); +#endif #define OMP_WAITING 5 #define OMP_RUNNING 6 From 2a5221032546a9a93bc05a1c46b5235a6507f3b4 Mon Sep 17 00:00:00 2001 From: Christian Stewart Date: Wed, 26 Aug 2026 07:58:56 -0700 Subject: [PATCH 2/3] fix(hip): link MIOpen for convolution builds Remove the duplicate convolution include guard and link the MIOpen library explicitly. This lets HIP environments compile when the fused trainer includes the convolution encoder source. Signed-off-by: Christian Stewart --- build-hip.sh | 2 +- src-hip/cudnn_conv2d.hip.cpp | 2 -- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/build-hip.sh b/build-hip.sh index 827bbc1c3c..d7ee6903bd 100755 --- a/build-hip.sh +++ b/build-hip.sh @@ -45,7 +45,7 @@ echo "Compiling HIP training backend ($ARCH)..." src-hip/cuda_shim.cpp -o build/cuda_shim.o g++ -shared -fPIC -fopenmp \ build/bindings_hip.o build/cuda_shim.o "$STATIC_LIB" "$RAYLIB_NAME/lib/libraylib.a" \ - -L$ROCM/lib -lamdhip64 -lhipblas -lhiprand -lrocrand -lrccl \ + -L$ROCM/lib -lamdhip64 -lhipblas -lhiprand -lrocrand -lrccl -lMIOpen \ -Wl,-rpath,$ROCM/lib -Bsymbolic-functions \ -lm -lpthread -lomp5 \ -o "$OUTPUT" diff --git a/src-hip/cudnn_conv2d.hip.cpp b/src-hip/cudnn_conv2d.hip.cpp index 1598ce89fa..4526c86995 100644 --- a/src-hip/cudnn_conv2d.hip.cpp +++ b/src-hip/cudnn_conv2d.hip.cpp @@ -4,8 +4,6 @@ // ocean.hip.cpp can include this file where ocean.cu includes the cuDNN // original. Included by ocean.hip.cpp (training). -#ifndef CUDNN_CONV2D_CU -#define CUDNN_CONV2D_CU #ifndef CUDNN_CONV2D_CU #define CUDNN_CONV2D_CU From cd53520b39ff34484763722ee7ebfc5f31bae86b Mon Sep 17 00:00:00 2001 From: Christian Stewart Date: Wed, 26 Aug 2026 22:39:46 -0700 Subject: [PATCH 3/3] fix(pufferl): load checkpoints before training Honor --load-model-path in train mode after the backend creates the policy and before the first rollout. Resolve latest checkpoints with the same directory convention as eval mode. Without this call, a requested resume silently started from random weights. Signed-off-by: Christian Stewart --- pufferlib/pufferl.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/pufferlib/pufferl.py b/pufferlib/pufferl.py index 929c300b48..c3d490620e 100644 --- a/pufferlib/pufferl.py +++ b/pufferlib/pufferl.py @@ -224,6 +224,21 @@ def _train(env_name, args, sweep_obj=None, result_queue=None, verbose=False): result_queue.put((args['gpu_id'], [], [], [])) return + load_path = args.get('load_model_path') + if load_path: + if load_path == 'latest': + pattern = os.path.join( + args['checkpoint_dir'], args['env_name'], '**', '*.bin') + candidates = glob.glob(pattern, recursive=True) + if not candidates: + raise FileNotFoundError( + f'No .bin checkpoints found in ' + f'{args["checkpoint_dir"]}/{args["env_name"]}/') + load_path = max(candidates, key=os.path.getctime) + backend.load_weights(pufferl, load_path) + if verbose: + print(f'Loaded weights from {load_path}') + args.pop('nccl_id', None) model_size = pufferl.num_params() if verbose: