From cadf2c1f093f6d1ad094f7236f7ebb05110f8411 Mon Sep 17 00:00:00 2001 From: Luna Nova Date: Sat, 3 Jan 2026 12:18:33 -0800 Subject: [PATCH] opensplat: add ROCm support --- pkgs/by-name/op/opensplat/package.nix | 57 +++++++++++++++++++++++++-- pkgs/top-level/all-packages.nix | 1 + 2 files changed, 54 insertions(+), 4 deletions(-) diff --git a/pkgs/by-name/op/opensplat/package.nix b/pkgs/by-name/op/opensplat/package.nix index 45791b230c1b..a4c417c9be8e 100644 --- a/pkgs/by-name/op/opensplat/package.nix +++ b/pkgs/by-name/op/opensplat/package.nix @@ -12,15 +12,16 @@ cxxopts, nix-update-script, config, - # Upstream has rocm/hip support, too. anyone? cudaSupport ? config.cudaSupport, cudaPackages, + rocmSupport ? config.rocmSupport, + rocmPackages, autoAddDriverRunpath, fetchpatch2, }: let version = "1.1.4"; - torch = python3.pkgs.torch.override { inherit cudaSupport; }; + torch = python3.pkgs.torch.override { inherit cudaSupport rocmSupport; }; # Using a normal stdenv with cuda torch gives # ld: /nix/store/k1l7y96gv0nc685cg7i3g43i4icmddzk-python3.11-torch-2.2.1-lib/lib/libc10.so: undefined reference to `std::ios_base_library_init()@GLIBCXX_3.4.32' stdenv' = if cudaSupport then cudaPackages.backendStdenv else stdenv; @@ -43,10 +44,32 @@ stdenv'.mkDerivation { }) ]; + postPatch = lib.optionalString rocmSupport '' + # ROCm CMake targets must be available before find_package(Torch) + # because Torch's Caffe2Targets.cmake references them in torch_hip_library + substituteInPlace CMakeLists.txt \ + --replace-fail "find_package(Torch REQUIRED)" \ + "find_package(hip REQUIRED) + find_package(hiprtc REQUIRED) + find_package(hipblas REQUIRED) + find_package(hipfft REQUIRED) + find_package(hiprand REQUIRED) + find_package(hipsparse REQUIRED) + find_package(hipsolver REQUIRED) + find_package(hipblaslt REQUIRED) + find_package(rocblas REQUIRED) + find_package(rocsolver REQUIRED) + find_package(miopen REQUIRED) + find_package(Torch REQUIRED)" + ''; + nativeBuildInputs = [ cmake ninja ] + ++ lib.optionals rocmSupport [ + rocmPackages.clr + ] ++ lib.optionals cudaSupport [ cudaPackages.cuda_nvcc autoAddDriverRunpath @@ -61,16 +84,39 @@ stdenv'.mkDerivation { torch opencv ] + ++ lib.optionals rocmSupport [ + rocmPackages.clr + rocmPackages.hipblas + rocmPackages.hipfft + rocmPackages.hiprand + rocmPackages.hipsparse + rocmPackages.hipsolver + rocmPackages.hipblaslt + rocmPackages.rocblas + rocmPackages.rocsolver + rocmPackages.rocsolver + rocmPackages.miopen + ] ++ lib.optionals cudaSupport [ cudaPackages.cuda_cudart ]; - env.TORCH_CUDA_ARCH_LIST = "${lib.concatStringsSep ";" python3.pkgs.torch.cudaCapabilities}"; + env = + lib.optionalAttrs cudaSupport { + TORCH_CUDA_ARCH_LIST = "${lib.concatStringsSep ";" python3.pkgs.torch.cudaCapabilities}"; + NIX_LDFLAGS = "-L${lib.getOutput "stubs" cudaPackages.cuda_cudart}/lib/stubs"; # fixes -lcuda not found + } + // lib.optionalAttrs rocmSupport { + HIPFLAGS = "-I${lib.getInclude rocmPackages.rocthrust}/include -I${lib.getInclude rocmPackages.rocprim}/include"; + }; cmakeFlags = [ (lib.cmakeBool "CMAKE_SKIP_RPATH" true) (lib.cmakeFeature "FETCHCONTENT_TRY_FIND_PACKAGE_MODE" "ALWAYS") ] + ++ lib.optionals rocmSupport [ + (lib.cmakeFeature "GPU_RUNTIME" "HIP") + ] ++ lib.optionals cudaSupport [ (lib.cmakeFeature "GPU_RUNTIME" "CUDA") (lib.cmakeFeature "CUDA_TOOLKIT_ROOT_DIR" "${cudaPackages.cudatoolkit}/") @@ -87,7 +133,10 @@ stdenv'.mkDerivation { # vendored+modified gsplat lib.licenses.asl20 ]; - maintainers = [ lib.maintainers.jcaesar ]; + maintainers = [ + lib.maintainers.jcaesar + lib.maintainers.LunNova + ]; platforms = lib.platforms.linux ++ lib.optionals (!cudaSupport) lib.platforms.darwin; }; } diff --git a/pkgs/top-level/all-packages.nix b/pkgs/top-level/all-packages.nix index 1b49fa2ec8d9..fc41a46d7d04 100644 --- a/pkgs/top-level/all-packages.nix +++ b/pkgs/top-level/all-packages.nix @@ -4104,6 +4104,7 @@ with pkgs; colmapWithCuda = colmap.override { cudaSupport = true; }; + opensplatWithRocm = opensplat.override { rocmSupport = true; }; opensplatWithCuda = opensplat.override { cudaSupport = true; }; chickenPackages_4 = recurseIntoAttrs (callPackage ../development/compilers/chicken/4 { });