GPU example
This page shows a GPU workflow for tractography. The implementation is backend-agnostic (CUDA, Metal, ...).
To maximize throughput, we use Float32.
using Tractography
const TG = Tractography
# use the Model as the streamline model
model = Model(Δt = 0.125f0,
foddata = FODData((@__DIR__) * "/../../examples/fod-FC.nii.gz"),
cone = Cone(45f0),
proba_min = 0.005f0,
)Model with elype Float32
├─ Δt = 0.125
├─ minimal probability = 0.005
├─ cone = Cone{Float32}(45.0f0)
├─ mollifier = max_mollifier
├─ evaluation of the basis = PreComputeAllFOD()
└─ data : (lmax = 8)
Define the seeds
using CUDA
# number of streamlines
Nmc = 1024 * 400
# maximum number of steps for each streamline
Nt = 2000
# define the seeds
seeds = cu(zeros(6, Nmc));
seeds[1:3, :] .= [-13.75, 26.5, 8] .+ 0.1 .* randn(3, Nmc) .|> Float32 |> CuArray;
seeds[4, :] .= 1
tract_length = CuArray(zeros(UInt32, Nmc))GPU on Apple OSX
using Metal
cu = MtlArray{Float32}
# number of streamlines
Nmc = 1024 * 400
# maximum number of steps for each streamline
Nt = 2000
# define the seeds
seeds = cu(zeros(6, Nmc));
seeds[1:3, :] .= [-13.75, 26.5, 8] .+ 0.1 .* randn(3, Nmc) .|> Float32 |> MtlArray;
seeds[4, :] .= 1
tract_length = MtlArray(zeros(UInt32, Nmc))The rest of the workflow is unchanged.
streamlines_gpu = cu(zeros(Float32, 3, Nt, Nmc), unified = true)Define the computation cache
Because we often run multiple batches on the same model, precomputing the cache is recommended.
# we precompute the cache which is heavy otherwise each call to sample
# will recompute it
cache_g = TG.init(model, Probabilistic();
𝒯ₐ = CuArray,
n_sphere = 400);Compute the streamlines
The following takes 0.5s on a A100.
# this setup works for a GPU with 40GiB
# it yields 1e6/sec streamlines for Probabilistic
# and 2.2e6/sec streamlines for Deterministic
CUDA.@time TG.sample!(
streamlines_gpu,
tract_length,
model,
cache_g,
Probabilistic(),
seeds;
gputhreads = 1024,
);streamlines_gpu stays on device memory. To inspect the result on CPU, copy or wrap depending on your backend.
streamlines = @time unsafe_wrap(Array, streamlines_gpu);