python3Packages.torch.tests.*compile*: init
This commit is contained in:
@@ -0,0 +1,38 @@
|
||||
{
|
||||
cudaPackages,
|
||||
feature ? null,
|
||||
lib,
|
||||
libraries,
|
||||
name ? if feature == null then "torch-compile-cpu" else "torch-compile-${feature}",
|
||||
pythonPackages,
|
||||
stdenv,
|
||||
}:
|
||||
let
|
||||
deviceStr = if feature == null then "" else '', device="cuda"'';
|
||||
in
|
||||
(cudaPackages.writeGpuTestPython.override { python3Packages = pythonPackages; })
|
||||
{
|
||||
inherit name feature libraries;
|
||||
makeWrapperArgs = [
|
||||
"--suffix"
|
||||
"PATH"
|
||||
":"
|
||||
"${lib.getBin stdenv.cc}/bin"
|
||||
];
|
||||
}
|
||||
''
|
||||
import torch
|
||||
|
||||
|
||||
@torch.compile
|
||||
def opt_foo2(x, y):
|
||||
a = torch.sin(x)
|
||||
b = torch.cos(y)
|
||||
return a + b
|
||||
|
||||
|
||||
print(
|
||||
opt_foo2(
|
||||
torch.randn(10, 10${deviceStr}),
|
||||
torch.randn(10, 10${deviceStr})))
|
||||
''
|
||||
@@ -1,6 +1,6 @@
|
||||
{ callPackage }:
|
||||
|
||||
{
|
||||
rec {
|
||||
# To perform the runtime check use either
|
||||
# `nix run .#python3Packages.torch.tests.tester-cudaAvailable` (outside the sandbox), or
|
||||
# `nix build .#python3Packages.torch.tests.tester-cudaAvailable.gpuCheck` (in a relaxed sandbox)
|
||||
@@ -14,4 +14,18 @@
|
||||
versionAttr = "hip";
|
||||
libraries = ps: [ ps.torchWithRocm ];
|
||||
};
|
||||
|
||||
compileCpu = tester-compileCpu.gpuCheck;
|
||||
tester-compileCpu = callPackage ./mk-torch-compile-check.nix {
|
||||
feature = null;
|
||||
libraries = ps: [ ps.torch ];
|
||||
};
|
||||
tester-compileCuda = callPackage ./mk-torch-compile-check.nix {
|
||||
feature = "cuda";
|
||||
libraries = ps: [ ps.torchWithCuda ];
|
||||
};
|
||||
tester-compileRocm = callPackage ./mk-torch-compile-check.nix {
|
||||
feature = "rocm";
|
||||
libraries = ps: [ ps.torchWithRocm ];
|
||||
};
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user