diff --git a/.buildkite/compare_regression.jl b/.buildkite/compare_regression.jl new file mode 100644 index 000000000..c85009bfd --- /dev/null +++ b/.buildkite/compare_regression.jl @@ -0,0 +1,116 @@ +# Compare benchmark harness results from the base branch and this commit. +# +# julia compare_regression.jl [--threshold=10] [--base=main] [--out=report.md] +# +# Each root holds the harness run directory for that side. Exits 1 on a +# slowdown above the threshold or a failed candidate run; base failures are +# only reported so a PR can fix them. + +using Printf +using Statistics +using TOML + +function load_runs(root) + runs = Dict{Tuple,Union{Float64,Nothing}}() + isdir(root) || return runs + for dir in readdir(root; join=true) + manifest_path = joinpath(dir, "manifest.toml") + isfile(manifest_path) || continue + for r in TOML.parsefile(manifest_path)["runs"] + key = (r["name"], r["T"], r["N"], r["M"]) + runs[key] = r["status"] == "complete" ? median_time(dir, r) : nothing + end + end + return runs +end + +# Median trial time in ms, or `nothing` when rows are missing or incorrect. +function median_time(dir, r) + # The harness writes `_.csv`, where save_as extends the model + # (e.g. `cunumeric_nofusion`, `cunumeric_struct`). + results = joinpath(dir, r["results_subdir"]) + isdir(results) || return nothing + prefix = "$(r["name"])_$(r["model"])" + files = filter(readdir(results)) do f + return f == "$prefix.csv" || (startswith(f, "$(prefix)_") && endswith(f, ".csv")) + end + times = Float64[] + for path in joinpath.(results, files), line in eachline(path) + f = split(strip(line), ',') + length(f) == 8 || continue + (parse(Int, f[3]), parse(Int, f[4])) == (r["N"], r["M"]) || continue + f[8] == "fail" && return nothing + push!(times, parse(Float64, f[6])) + end + return isempty(times) ? nothing : median(times) +end + +label(key) = "$(key[1]) ($(key[2]), $(key[3])×$(key[4]))" + +function report(base_root, candidate_root, threshold, base_name) + base = load_runs(base_root) + candidate = load_runs(candidate_root) + isempty(candidate) && error("no candidate results in $candidate_root") + + lines = [ + "## Benchmark regression vs. `$base_name`", "", + "| Benchmark | `$base_name` (ms) | PR (ms) | Change |", + "| --- | ---: | ---: | ---: |", + ] + regressions, failed, uncompared = String[], String[], String[] + for key in sort!(collect(keys(candidate))) + after = candidate[key] + before = get(base, key, nothing) + if after === nothing + push!(failed, label(key)) + elseif before === nothing + push!(uncompared, label(key)) + else + change = 100 * (after / before - 1) + flag = change > threshold ? " ⚠️" : "" + push!( + lines, + @sprintf("| %s | %.3f | %.3f | %+.1f%%%s |", + label(key), before, after, change, flag) + ) + change > threshold && + push!(regressions, @sprintf("%s: %.1f%% slower", label(key), change)) + end + end + for (title, items) in (("Slower than the threshold", regressions), + ("Failed on this PR", failed), + ("Not compared (no base result)", uncompared)) + isempty(items) && continue + append!(lines, ["", "**$title:**"], ["- $item" for item in items]) + end + push!(lines, "", @sprintf("Threshold: more than %.0f%% slower (median of trials).", threshold)) + return join(lines, '\n') * '\n', isempty(regressions) && isempty(failed) +end + +function main(args) + threshold, out, base_name = 10.0, nothing, "base" + positional = String[] + for arg in args + if startswith(arg, "--threshold=") + threshold = parse(Float64, split(arg, '='; limit=2)[2]) + elseif startswith(arg, "--base=") + base_name = split(arg, '='; limit=2)[2] + elseif startswith(arg, "--out=") + out = split(arg, '='; limit=2)[2] + else + push!(positional, arg) + end + end + length(positional) == 2 || error("usage: compare_regression.jl ") + text, ok = report(positional..., threshold, base_name) + print(text) + out === nothing || write(out, text) + return ok ? 0 : 1 +end + +try + exit(main(ARGS)) +catch e + println(stderr, "comparison failed: ", sprint(showerror, e)) + exit(2) +end diff --git a/.buildkite/install_cmake.sh b/.buildkite/install_cmake.sh new file mode 100644 index 000000000..78ca58bb6 --- /dev/null +++ b/.buildkite/install_cmake.sh @@ -0,0 +1,11 @@ +# shellcheck shell=bash +# Source to put a pinned CMake first on PATH (developer wrapper builds need it). +CMAKE_VERSION="3.30.7" +CMAKE_ROOT="$(mktemp -d)" +CMAKE_INSTALLER="$CMAKE_ROOT/cmake-installer.sh" +curl --fail --silent --show-error --location \ + --output "$CMAKE_INSTALLER" \ + "https://github.com/Kitware/CMake/releases/download/v$CMAKE_VERSION/cmake-$CMAKE_VERSION-linux-x86_64.sh" +sh "$CMAKE_INSTALLER" --skip-license --prefix="$CMAKE_ROOT" +export PATH="$CMAKE_ROOT/bin:$PATH" +cmake --version diff --git a/.buildkite/regression.pipeline.yml b/.buildkite/regression.pipeline.yml new file mode 100644 index 000000000..8298f5f37 --- /dev/null +++ b/.buildkite/regression.pipeline.yml @@ -0,0 +1,21 @@ +steps: + - label: ":chart_with_upwards_trend: Benchmark regression vs. base branch" + key: "regression" + plugins: + - JuliaCI/julia#v1: + version: "1.12" + cache_dir: "${HOME}/.cache/julia-buildkite-plugin-regression" + command: ".buildkite/run_regression.sh" + artifact_paths: + - "regression/**/*" + agents: + queue: "cuda" + # One comparison at a time so runs do not share a GPU. + concurrency: 1 + concurrency_group: "cunumeric/regression" + timeout_in_minutes: 180 + env: + LD_LIBRARY_PATH: "" + LEGATE_AUTO_CONFIG: "0" + # nvidia-smi cannot read GPU memory in the CI container; cap the harness budget. + CUNUMERIC_BENCH_FBMEM_MB: "3072" diff --git a/.buildkite/regression.toml b/.buildkite/regression.toml new file mode 100644 index 000000000..fdc61a476 --- /dev/null +++ b/.buildkite/regression.toml @@ -0,0 +1,86 @@ +# Opt-in performance regression suite ([regression-ci]): cuNumeric on the base +# branch vs. this commit. Sizes are pinned so both sides run identical problems. +# gemm, montecarlo(_naive), grayscott_plain and cg_plain avoid @accelerate and +# newer APIs, so they are the entries comparable against main. +[Global] +models = ["cunumeric"] +n_warmup = 2 +n_iter = 5 +n_trial = 5 +check_correctness = true +auto_size = false +# Legate's pool is capped at CUNUMERIC_BENCH_FBMEM_MB (3 GB) on CI. +mem_frac = 0.9 +# cpus = 4: 8 CPU procs + GPU/util threads exceed the CI agent's cores. + +[[gemm]] +T = "Float32" +N = 4096 +M = 4096 +gpus = 1 +cpus = 4 + +[[montecarlo]] +T = "Float32" +N = 50_000_000 +M = 1 +gpus = 1 +cpus = 4 + +[[montecarlo_naive]] +T = "Float32" +N = 50_000_000 +M = 1 +gpus = 1 +cpus = 4 + +[[grayscott]] +T = "Float32" +N = 1024 +M = 1024 +gpus = 1 +cpus = 4 + +[[grayscott_plain]] +T = "Float32" +N = 1024 +M = 1024 +gpus = 1 +cpus = 4 + +[[cg]] +T = "Float64" +N = 65536 +M = 1 +gpus = 1 +cpus = 4 +kwargs = { check_every = 10, max_iter = 1000 } + +[[cg_plain]] +T = "Float64" +N = 65536 +M = 1 +gpus = 1 +cpus = 4 +kwargs = { check_every = 10, max_iter = 1000 } + +[[nas_ep]] +T = "Float64" +n_iter = 1 +gpus = 1 +cpus = 4 +kwargs = { class = "A" } + +[[nas_mg]] +T = "Float64" +n_iter = 1 +gpus = 1 +cpus = 4 +kwargs = { class = "A" } + +[[nas_ft]] +T = "Float64" +n_iter = 1 +gpus = 1 +cpus = 4 +kwargs = { class = "A" } diff --git a/.buildkite/regression_base_pins.toml b/.buildkite/regression_base_pins.toml new file mode 100644 index 000000000..6b1928aa1 --- /dev/null +++ b/.buildkite/regression_base_pins.toml @@ -0,0 +1,6 @@ +# Extra version pins for the base side's harness environment, keyed by base +# branch, for bases that no longer load against the current registry. + +[main] +# cuNumeric 0.2.0 calls LegatePreferences.has_cuda_gpu, removed in 0.1.7. +LegatePreferences = "0.1.6" diff --git a/.buildkite/run_developer_ci.sh b/.buildkite/run_developer_ci.sh index f9af18261..856f758c2 100755 --- a/.buildkite/run_developer_ci.sh +++ b/.buildkite/run_developer_ci.sh @@ -10,15 +10,7 @@ case "${CUNUMERIC_FUSION:-}" in ;; esac -CMAKE_VERSION="3.30.7" -CMAKE_ROOT="$(mktemp -d)" -CMAKE_INSTALLER="$CMAKE_ROOT/cmake-installer.sh" -curl --fail --silent --show-error --location \ - --output "$CMAKE_INSTALLER" \ - "https://github.com/Kitware/CMake/releases/download/v$CMAKE_VERSION/cmake-$CMAKE_VERSION-linux-x86_64.sh" -sh "$CMAKE_INSTALLER" --skip-license --prefix="$CMAKE_ROOT" -export PATH="$CMAKE_ROOT/bin:$PATH" -cmake --version +source .buildkite/install_cmake.sh # Exercise libcxxwrap cache validation separately from package tests, which run in # JLL jobs where no build toolchain is installed. diff --git a/.buildkite/run_regression.sh b/.buildkite/run_regression.sh new file mode 100755 index 000000000..33f213b0e --- /dev/null +++ b/.buildkite/run_regression.sh @@ -0,0 +1,119 @@ +#!/usr/bin/env bash +# Opt-in GPU performance comparison of this commit against a base branch: +# [regression-ci ], else REGRESSION_BASE_BRANCH, else the branch the PR +# targets, else main. Both sides run this commit's harness and regression.toml. + +set -euo pipefail + +requested="$(buildkite-agent meta-data get regression-base-branch --default "" 2>/dev/null || true)" +pr_base="${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-}" +readonly BASE_BRANCH="${requested:-${REGRESSION_BASE_BRANCH:-${pr_base:-main}}}" +readonly THRESHOLD="${REGRESSION_THRESHOLD:-10}" +# Bounds each base harness pass so a hang on the base cannot use up the step. +readonly BASE_TIMEOUT="${REGRESSION_BASE_TIMEOUT:-3600}" + +candidate="$PWD" +harness="$candidate/benchmark" +env_dir="$harness/environments/cunumeric" +out="$candidate/regression" +base="$(mktemp -d)/base" + +rm -rf "$out" +mkdir -p "$out" + +git fetch --no-tags origin "+refs/heads/$BASE_BRANCH:refs/remotes/origin/$BASE_BRANCH" +git worktree add --detach "$base" "origin/$BASE_BRANCH" +trap 'git -C "$candidate" worktree remove --force "$base"' EXIT +git submodule update --init benchmark + +julia --color=yes --project="$harness" -e 'using Pkg; Pkg.instantiate()' +mkdir -p "$harness/results" + +# Point the harness's cuNumeric environment at one checkout. Wrapper overrides +# and precompiled wrapper bindings from the other side must not leak across. +bind_checkout() { + local source=$1 mode=$2 pins=${3:-} + local depot + depot="$(julia --startup-file=no -e 'print(DEPOT_PATH[1])')" + rm -rf "$depot"/packages/*/*/override \ + "$depot"/compiled/v*/{cuNumeric,Legate,cunumeric_jl_wrapper_jll,legate_jl_wrapper_jll} + rm -f "$env_dir/Manifest.toml" "$env_dir/LocalPreferences.toml" + # The harness pins the current cuNumeric/CNPreferences; each side develops + # its own checkout, so drop those pins to let an older base resolve. + julia --color=yes --project="$env_dir" -e ' + using Pkg, TOML + project = Base.active_project() + toml = TOML.parsefile(project) + foreach(p -> delete!(get(toml, "compat", Dict()), p), ("cuNumeric", "CNPreferences")) + open(io -> TOML.print(io, toml), project, "w") + Pkg.develop([PackageSpec(path = ARGS[1]), PackageSpec(path = ARGS[2])]) + pins = isempty(ARGS[4]) ? Dict() : get(TOML.parsefile(ARGS[3]), ARGS[4], Dict()) + isempty(pins) || Pkg.add([PackageSpec(name = k, version = v) for (k, v) in pins]) + Pkg.instantiate() + ' "$source" "$source/lib/CNPreferences" "$candidate/.buildkite/regression_base_pins.toml" "$pins" + if [[ "$mode" == developer ]]; then + julia --color=yes --project="$env_dir" -e ' + using CNPreferences, Pkg + CNPreferences.use_developer_mode() + Pkg.build("cuNumeric") + ' + fi +} + +result_dirs() { find "$harness/results" -mindepth 1 -maxdepth 1 -type d -printf '%f\n'; } + +run_side() { + local side=$1 source=$2 mode=$3 + local limit=() + [[ "$side" == base ]] && limit=(timeout --signal=KILL "$BASE_TIMEOUT") + echo "--- :julia: $side ($mode wrapper)" + bind_checkout "$source" "$mode" "$([[ "$side" == base ]] && echo "$BASE_BRANCH")" + local before new + before="$(result_dirs)" + (cd "$harness" && "${limit[@]}" julia --color=yes --project=. run.jl \ + --config="$candidate/.buildkite/regression.toml" --fusion=on) || + echo "Harness reported failures ($side)." + new="$(comm -13 <(sort <<<"$before") <(result_dirs | sort) | head -1)" + if [[ -n "$new" ]]; then + mkdir -p "$out/$side" + mv "$harness/results/$new" "$out/$side/results" + fi +} + +# Build a side's wrapper from source when it differs from the release its own +# checkout records. A base without RELEASED_COMMIT is checked against ours. +wrapper_mode() { + local dir=$1 + [[ -f "$dir/scripts/wrapper_changed.sh" ]] || dir="$candidate" + if (cd "$dir" && scripts/wrapper_changed.sh "$2" >&2); then + echo jll + else + local status=$? + ((status == 1)) || exit "$status" + echo developer + fi +} + +base_mode="$(wrapper_mode "$base" "$(git -C "$base" rev-parse HEAD)")" || base_mode=developer +candidate_mode="$(wrapper_mode "$candidate" HEAD)" +if [[ "$base_mode" == developer || "$candidate_mode" == developer ]]; then + source .buildkite/install_cmake.sh +fi + +echo "Comparing against $BASE_BRANCH." +# Base failures only leave its results uncompared; the candidate must pass. +run_side base "$base" "$base_mode" || echo "Base side failed; its results are not compared." +run_side candidate "$candidate" "$candidate_mode" + +echo "--- :bar_chart: Compare" +status=0 +julia --startup-file=no "$candidate/.buildkite/compare_regression.jl" \ + "$out/base" "$out/candidate" --threshold="$THRESHOLD" --base="$BASE_BRANCH" \ + --out="$out/report.md" || + status=$? + +if [[ -f "$out/report.md" ]] && command -v buildkite-agent >/dev/null; then + style=$([[ $status == 0 ]] && echo success || echo error) + buildkite-agent annotate --context regression --style "$style" < "$out/report.md" +fi +exit "$status" diff --git a/.buildkite/upload_gpu_ci.sh b/.buildkite/upload_gpu_ci.sh index 2558339dc..abbe8b460 100755 --- a/.buildkite/upload_gpu_ci.sh +++ b/.buildkite/upload_gpu_ci.sh @@ -7,6 +7,7 @@ readonly DEVELOPER_PIPELINE=".buildkite/developer.pipeline.yml" branch="${BUILDKITE_BRANCH:-}" base_branch="${BUILDKITE_PULL_REQUEST_BASE_BRANCH:-}" +pull_request="${BUILDKITE_PULL_REQUEST:-false}" message="${BUILDKITE_MESSAGE:-}" run_jll=true @@ -44,6 +45,27 @@ if [[ "$branch" != "main" && "$base_branch" != "main" ]]; then fi fi +# Opt-in performance comparison: [regression-ci] in the commit message or the +# pull request title/body compares against the PR's base branch, and +# [regression-ci ] against . +opt_in="$message" +if [[ "$pull_request" =~ ^[0-9]+$ ]]; then + opt_in+=$'\n'"$( + curl --fail --silent --show-error --location \ + --header "Accept: application/vnd.github+json" \ + "https://api.github.com/repos/JuliaLegate/cuNumeric.jl/pulls/$pull_request" | + python3 -c 'import json, sys; pr = json.load(sys.stdin); print(pr.get("title") or "", pr.get("body") or "")' + )" || true +fi +if [[ "$opt_in" =~ \[regression-ci([[:space:]]+([A-Za-z0-9._/-]+))?\] ]]; then + regression_base="${BASH_REMATCH[2]}" + if [[ -n "$regression_base" ]]; then + buildkite-agent meta-data set regression-base-branch "$regression_base" + fi + echo "Uploading benchmark regression CI (base: ${regression_base:-PR base branch})." + buildkite-agent pipeline upload .buildkite/regression.pipeline.yml +fi + # Each dynamic upload is inserted immediately after this job, so upload the # developer group first to keep the JLL group first when both suites run. if [[ "$run_developer" == "true" ]]; then diff --git a/scripts/wrapper_changed.sh b/scripts/wrapper_changed.sh index c9f39ba0d..0b3edc67b 100755 --- a/scripts/wrapper_changed.sh +++ b/scripts/wrapper_changed.sh @@ -1,15 +1,16 @@ #!/usr/bin/env bash -# Exit 0 if the wrapper at HEAD matches RELEASED_COMMIT (the source of the -# released wrapper JLL), 1 if it differs. +# Exit 0 if the wrapper at REF (default HEAD) matches RELEASED_COMMIT (the +# source of the released wrapper JLL), 1 if it differs. set -euo pipefail readonly WRAPPER_PATH="lib/cunumeric_jl_wrapper" readonly RELEASED_COMMIT_FILE="$WRAPPER_PATH/RELEASED_COMMIT" +ref="${1:-HEAD}" released="$(tr -d '[:space:]' < "$RELEASED_COMMIT_FILE")" if ! git cat-file -e "${released}^{commit}" 2>/dev/null; then git fetch --no-tags --depth=1 origin "$released" fi -git diff --quiet "$released" HEAD -- "$WRAPPER_PATH" ":(exclude)$RELEASED_COMMIT_FILE" +git diff --quiet "$released" "$ref" -- "$WRAPPER_PATH" ":(exclude)$RELEASED_COMMIT_FILE" diff --git a/src/ndarray/broadcast.jl b/src/ndarray/broadcast.jl index 8bf5029ca..bb6812bf8 100644 --- a/src/ndarray/broadcast.jl +++ b/src/ndarray/broadcast.jl @@ -148,15 +148,18 @@ function __materialize(bc::Broadcasted{<:NDArrayStyle}) end # The C API is binary, so evaluate flattened `+` and `*` chains pairwise. -function _unravel_flattened_associative(f, args::Tuple) +function _unravel_flattened_associative(f, args::Tuple, dest) acc = first(args) owns_acc = false - for arg in Base.tail(args) - next = try - __materialize(Base.broadcasted(f, acc, arg)) - finally - owns_acc && acc isa NDArray && destroy!(acc) + for (i, arg) in enumerate(Base.tail(args)) + bc = Base.broadcasted(f, acc, arg) + # Only the final binary operation may overwrite the destination. + next = if i == length(args) - 1 && !isnothing(dest) + unravel_broadcast_tree(Base.Broadcast.instantiate(bc), dest) + else + __materialize(bc) end + owns_acc && acc isa NDArray && destroy!(acc) acc = next owns_acc = acc isa NDArray end @@ -178,7 +181,7 @@ end # top-level result directly if its eltype matches and no input partially overlaps it. function unravel_broadcast_tree(bc::Broadcasted, dest=nothing) if length(bc.args) > 2 && _is_flattened_associative(bc.f) - return _unravel_flattened_associative(bc.f, bc.args) + return _unravel_flattened_associative(bc.f, bc.args, dest) end # Recursively materialize/unravel any nested broadcasts @@ -224,16 +227,12 @@ end return result === dest ? dest : _copyto_unfused!(dest, result) end -# Slice destinations must assign into their parent store. +# Preserve the destination store: other handles may already view it. @inline function _store_broadcast_result!( dest::NDArray{T}, temp_result::NDArray{T} ) where {T} - if _is_ndarray_slice(dest) - nda_assign(dest, temp_result) - destroy!(temp_result) - else - nda_move(dest, temp_result) - end + nda_assign(dest, temp_result) + destroy!(temp_result) return dest end @@ -320,11 +319,11 @@ end return fuse_broadcast_tree!(dest, bc) else _assert_struct_broadcast_fused(dest, bc) - return _copyto_unfused!(dest, unravel_broadcast_tree(bc)) + return _unfused_into!(dest, bc) end else _assert_struct_broadcast_fused(dest, bc) - return _copyto_unfused!(dest, unravel_broadcast_tree(bc)) + return _unfused_into!(dest, bc) end end diff --git a/src/ndarray/detail/ndarray.jl b/src/ndarray/detail/ndarray.jl index 46e6be468..a449d9ea7 100644 --- a/src/ndarray/detail/ndarray.jl +++ b/src/ndarray/detail/ndarray.jl @@ -843,8 +843,11 @@ end Return the size of the given `NDArray`. """ function shape(arr::NDArray{<:Any,N}) where {N} - shp = cuNumeric.nda_array_shape(arr) - return ntuple(i -> Int(shp[i]), Val(N)) + # Rank is known from the type; avoid a rank query and temporary shape Vector. + shp = Ref{NTuple{N,UInt64}}() + ccall((:nda_array_shape, libnda), + Cvoid, (NDArray_t, Ref{NTuple{N,UInt64}}), arr.ptr, shp) + return map(Int, shp[]) end @doc""" diff --git a/src/ndarray/diagonal.jl b/src/ndarray/diagonal.jl index fb06ef06d..19099d946 100644 --- a/src/ndarray/diagonal.jl +++ b/src/ndarray/diagonal.jl @@ -130,19 +130,35 @@ function Base.:*(D::DiagonalNDArray, v::NDArray{<:Any,1}) end function LinearAlgebra.lmul!(D::DiagonalNDArray, B::NDArray) - return copyto!(B, D * B) + return mul!(B, D, B) end function LinearAlgebra.rmul!(A::NDArray, D::DiagonalNDArray) - return copyto!(A, A * D) + return mul!(A, A, D) end -function LinearAlgebra.mul!(C::NDArray, D::DiagonalNDArray, A::NDArray) - return copyto!(C, D * A) -end - -function LinearAlgebra.mul!(C::NDArray, A::NDArray, D::DiagonalNDArray) - return copyto!(C, A * D) +function LinearAlgebra.mul!( + C::NDArray, D::DiagonalNDArray, A::Union{NDArray{<:Any,1},NDArray{<:Any,2}} +) + size(C) == size(A) || + throw(DimensionMismatch("diagonal product destination must have size $(size(A))")) + size(A, 1) == size(D, 1) || + throw(DimensionMismatch("diagonal product dimensions do not match")) + d = ndims(A) == 1 ? _diag_vec(D) : _row_scale(_diag_vec(D)) + C .= d .* A + d === _diag_vec(D) || destroy!(d) + return C +end + +function LinearAlgebra.mul!(C::NDArray, A::NDArray{<:Any,2}, D::DiagonalNDArray) + size(C) == size(A) || + throw(DimensionMismatch("diagonal product destination must have size $(size(A))")) + size(A, 2) == size(D, 1) || + throw(DimensionMismatch("diagonal product dimensions do not match")) + d = _col_scale(_diag_vec(D)) + C .= A .* d + destroy!(d) + return C end function Base.:\(D::DiagonalNDArray, B::NDArray{<:Any,1}) @@ -158,7 +174,10 @@ function Base.:\(D::DiagonalNDArray, B::NDArray{<:Any,2}) "matrix is $(size(B,1))×$(size(B,2)), but diagonal is $(size(D,1))×$(size(D,2))" ), ) - return B ./ _row_scale(_diag_vec(D)) + d = _row_scale(_diag_vec(D)) + result = B ./ d + destroy!(d) + return result end function Base.:/(A::NDArray{<:Any,2}, D::DiagonalNDArray) @@ -167,15 +186,35 @@ function Base.:/(A::NDArray{<:Any,2}, D::DiagonalNDArray) "matrix is $(size(A,1))×$(size(A,2)), but diagonal is $(size(D,1))×$(size(D,2))" ), ) - return A * inv(D) + d = _col_scale(_diag_vec(D)) + result = A ./ d + destroy!(d) + return result end -function LinearAlgebra.ldiv!(D::DiagonalNDArray, B::NDArray) - return copyto!(B, D \ B) +function LinearAlgebra.ldiv!(D::DiagonalNDArray, B::NDArray{<:Any,1}) + length(B) == size(D, 1) || + throw(DimensionMismatch("vector length $(length(B)) does not match diagonal $(size(D,1))")) + B ./= _diag_vec(D) + return B end -function LinearAlgebra.rdiv!(A::NDArray, D::DiagonalNDArray) - return copyto!(A, A / D) +function LinearAlgebra.ldiv!(D::DiagonalNDArray, B::NDArray{<:Any,2}) + size(B, 1) == size(D, 1) || + throw(DimensionMismatch("diagonal division dimensions do not match")) + d = _row_scale(_diag_vec(D)) + B ./= d + destroy!(d) + return B +end + +function LinearAlgebra.rdiv!(A::NDArray{<:Any,2}, D::DiagonalNDArray) + size(A, 2) == size(D, 1) || + throw(DimensionMismatch("diagonal division dimensions do not match")) + d = _col_scale(_diag_vec(D)) + A ./= d + destroy!(d) + return A end function Base.inv(D::DiagonalNDArray{T}) where {T} @@ -188,13 +227,17 @@ LinearAlgebra.det(D::DiagonalNDArray) = prod(_diag_vec(D)) LinearAlgebra.tr(D::DiagonalNDArray{<:Number}) = sum(_diag_vec(D)) Base.sum(D::DiagonalNDArray) = sum(_diag_vec(D)) -# Generic `prod(::AbstractMatrix)` walks every entry (scalar-indexing). For n>1 a -# Diagonal has off-diagonal zeros, so the product is zero — match Base, as 0D. +# Include structural zeros before reducing so nonfinite entries propagate, +# without overflowing a product of finite diagonal entries first. function Base.prod(D::DiagonalNDArray{T}) where {T<:Number} n = size(D, 1) n == 0 && return cnscalar(NDArray(one(T))) - n == 1 && return prod(_diag_vec(D)) - return cnscalar(NDArray(zero(T))) + n == 1 && return sum(_diag_vec(D)) + values = zero(T) .* _diag_vec(D) + # All finite terms are zero; sum also supports complex backend types. + result = sum(values) + destroy!(values) + return result end function Base.maximum(D::DiagonalNDArray{T}) where {T<:Number} @@ -286,19 +329,40 @@ end function LinearAlgebra.norm(D::DiagonalNDArray, p::Real=2) p = _maybe_fetch(p) - # Off-diagonals are zero, so the matrix vec-norm equals the diag vec-norm. + R = float(real(eltype(D))) + isempty(D) && return cnscalar(NDArray(zero(R))) d = abs.(_diag_vec(D)) - if p == 2 - return sqrt(sum(d .^ 2)) - elseif p == 1 - return sum(d) - elseif p == Inf - return isempty(D) ? cnscalar(NDArray(float(real(zero(eltype(D)))))) : maximum(d) - elseif p == -Inf - return isempty(D) ? cnscalar(NDArray(float(real(zero(eltype(D)))))) : minimum(d) - else - return sum(d .^ p) ^ (one(p) / p) - end + # Canonicalize host orders, including signed zero, to share specializations. + result = _diagonal_norm(d, Val(Float64(p) + 0.0), R) + destroy!(d) + return result +end + +function _diagonal_norm(d, ::Val{0.0}, ::Type{R}) where {R} + mask = d .!= zero(eltype(d)) + values = as_type(mask, R) + result = sum(values) + destroy!(values) + destroy!(mask) + return result +end + +_diagonal_norm(d, ::Val{1.0}, ::Type) = sum(d) +_diagonal_norm(d, ::Val{2.0}, ::Type) = sqrt(sum(d .^ 2)) +_diagonal_norm(d, ::Val{Inf}, ::Type) = maximum(d) + +function _diagonal_norm(d, ::Val{-Inf}, ::Type{R}) where {R} + # With structural zeros, every negative norm is zero (or NaN). Reuse the + # -1 case to propagate NaNs, which the backend minimum would discard. + return length(d) > 1 ? _diagonal_norm(d, Val(-1.0), R) : minimum(d) +end + +function _diagonal_norm(d, ::Val{P}, ::Type{R}) where {P,R} + total = sum(d .^ R(P)) + # Each structural zero contributes Inf for a finite negative p. + # Adding that contribution retains NaNs from diagonal entries. + length(d) > 1 && P < 0 && (total = total + R(Inf)) + return total ^ inv(R(P)) end function LinearAlgebra.cond(D::DiagonalNDArray, p::Real=2) @@ -308,7 +372,9 @@ function LinearAlgebra.cond(D::DiagonalNDArray, p::Real=2) end isempty(D) && return cnscalar(NDArray(float(one(real(eltype(D)))))) dabs = abs.(_diag_vec(D)) - return maximum(dabs) / minimum(dabs) + result = maximum(dabs) / minimum(dabs) + destroy!(dabs) + return result end function Base.:+(A::NDArray{T,2}, D::DiagonalNDArray) where {T} diff --git a/src/ndarray/ndarray.jl b/src/ndarray/ndarray.jl index 0057c629c..9cfaf2d55 100644 --- a/src/ndarray/ndarray.jl +++ b/src/ndarray/ndarray.jl @@ -162,6 +162,22 @@ function Base.copyto!(dest::NDArray{T,N}, src::Array{T,N}) where {T,N} return dest end +# Borrow a host vector only until the copy completes; avoid the constructor's +# separate owned copy and fence. Struct packing keeps its existing path. +function Base.copyto!(dest::NDArray{T,1}, src::Vector{T}) where {T<:SUPPORTED_ARRAY_TYPES} + size(dest) == size(src) || throw(DimensionMismatch("source and destination sizes differ")) + isempty(src) && return dest + GC.@preserve src begin + attached = nda_attach_external(src) + GC.@preserve attached dest begin + copyto!(dest, attached) + issue_execution_fence(; block=true) + end + destroy!(attached) + end + return dest +end + @doc""" as_type(arr::NDArray, t::Type{T}) where {T} @@ -424,6 +440,13 @@ Base.IndexStyle(::Type{<:NDArray}) = IndexCartesian() Base.axes(arr::NDArray) = Base.OneTo.(size(arr)) Base.view(arr::NDArray, inds...) = arr[inds...] # NDArray slices are views by default. +# All-colon getindex copies, but view/dotview must retain the original store. +# A separate handle is required because @accelerate frees temporary views. +function Base.view(arr::NDArray{T,N}, ::Vararg{Colon,N}) where {T,N} + return nda_get_slice(arr, slice_array((nothing, nothing))) +end +Base.view(arr::NDArray{T,0}) where {T} = nda_reshape_array(arr, ()) + function Base.show(io::IO, arr::NDArray{T,0}) where {T} print(io, summary(arr), "(") @allowscalar show(io, arr[]) @@ -462,7 +485,10 @@ end Overloads `Base.getindex` and `Base.setindex!` to support multidimensional indexing and slicing on `cuNumeric.NDArray`s. Slicing supports combinations of `Integer`, `UnitRange`, and `Colon()` for selecting ranges of rows and columns. -The use of all colons (`arr[:]`, `arr[:, :]`, etc.) returns a new Julia `Array` containing a copy of the data. +Using one colon per dimension (`v[:]`, `A[:, :]`, etc.) returns an `NDArray` +copy. `view` and dotted assignment share the original storage instead. +Linear range/colon indexing of multidimensional arrays is unsupported; use one +index per dimension. Assignment also supports: - Writing NDArray slices to NDArray regions @@ -751,26 +777,52 @@ end end @inline function Base.getindex( - arr::NDArray, i::AbstractUnitRange{<:Integer} -) + arr::NDArray{T,1}, i::AbstractUnitRange{<:Integer} +) where {T} @boundscheck checkbounds(arr, i) return nda_get_slice(arr, slice_array(_zero_based_range(i))) end +function Base.getindex(arr::NDArray, i::AbstractUnitRange{<:Integer}) + throw(ArgumentError( + "linear range indexing of multidimensional NDArrays is unsupported; " * + "use one index per dimension" + )) +end + +function Base.setindex!(arr::NDArray, rhs::NDArray, i::AbstractUnitRange{<:Integer}) + throw(ArgumentError( + "linear range assignment to multidimensional NDArrays is unsupported; " * + "use one index per dimension" + )) +end + @inline function Base.getindex( - arr::NDArray{T}, c::Vararg{Colon,N} + arr::NDArray{T,N}, c::Vararg{Colon,N} ) where {T,N} - @boundscheck checkbounds(arr, c...) return Base.copy(arr) end +function Base.getindex(arr::NDArray, ::Vararg{Colon}) + throw(ArgumentError( + "linear colon indexing of multidimensional NDArrays is unsupported; " * + "use one colon per dimension" + )) +end + @inline function Base.setindex!( - arr::NDArray{T}, rhs::NDArray{T}, c::Vararg{Colon,N} + arr::NDArray{T,N}, rhs::NDArray{T}, c::Vararg{Colon,N} ) where {T,N} - @boundscheck checkbounds(arr, c...) return Base.copyto!(arr, rhs) end +function Base.setindex!(arr::NDArray{T}, rhs::NDArray{T}, ::Vararg{Colon}) where {T} + throw(ArgumentError( + "linear colon assignment to multidimensional NDArrays is unsupported; " * + "use one colon per dimension" + )) +end + @inline function Base.setindex!( arr::NDArray{T,2}, val::T, ::Colon, j::Integer ) where {T} diff --git a/src/ndarray/sort.jl b/src/ndarray/sort.jl index c1c91fe62..39c3f992c 100644 --- a/src/ndarray/sort.jl +++ b/src/ndarray/sort.jl @@ -92,7 +92,7 @@ function searchsortedfirst(a::NDArray{T,1}, v::NDArray) where {T} end function searchsortedfirst(a::NDArray{T,1}, x::Number) where {T} - needle = NDArray(convert(T, x)) + needle = NDArray(x) result = searchsortedfirst(a, needle) destroy!(needle) return result @@ -103,7 +103,7 @@ function searchsortedlast(a::NDArray{T,1}, v::NDArray) where {T} end function searchsortedlast(a::NDArray{T,1}, x::Number) where {T} - needle = NDArray(convert(T, x)) + needle = NDArray(x) result = searchsortedlast(a, needle) destroy!(needle) return result diff --git a/src/scoping/accelerate.jl b/src/scoping/accelerate.jl index c7cb37aa1..2ddd6c47a 100644 --- a/src/scoping/accelerate.jl +++ b/src/scoping/accelerate.jl @@ -215,7 +215,14 @@ end Optimize straight-line array code by coordinating CUDA broadcast fusion within expressions, fusion across broadcast statements, and scope-aware cleanup of materialized temporaries. Control flow and nested/anonymous functions are -rejected. Four forms determine which values must remain valid: +rejected. + +Single-use producers may fuse across intervening read-only calculations when +their inputs are unchanged. Unknown calls remain barriers. Function-form fusion +across writes to other array arguments checks storage overlap at runtime and +retains materialized intermediates when those arguments overlap. + +Four forms determine which values must remain valid: * **function** (preferred): arguments and returned values are protected; non-returned locals may fuse into consumers or be freed after their last use. diff --git a/src/scoping/broadcast_lifetimes.jl b/src/scoping/broadcast_lifetimes.jl index e5568b5a4..49ae1c63e 100644 --- a/src/scoping/broadcast_lifetimes.jl +++ b/src/scoping/broadcast_lifetimes.jl @@ -116,11 +116,30 @@ function rewrite_broadcast_lifetimes(scope) return _prepend_statements(rewritten, temps), assigned_vars end +# Scalars/immutable scalar parameter records cannot share mutable array storage. +_fusion_disjoint(a, b) = isbitstype(typeof(a)) || isbitstype(typeof(b)) +_fusion_disjoint(a::AbstractArray, b::AbstractArray) = !Base.mightalias(a, b) +_fusion_disjoint(a::NDArray, b::AbstractArray) = false +_fusion_disjoint(a::AbstractArray, b::NDArray) = false +_fusion_disjoint(a::NDArray, b::NDArray) = !nda_overlaps(a, b) + function process_broadcast_lifetime_scope( scope; on_rewrite=nothing, protected_roots=Set{Symbol}() ) # Returned producers and caller-owned roots stay materialized: exempt from fusion. protected = union(_returned_symbols(scope), protected_roots) - scope = InterBroadcastFusion.rewrite_scope(scope; on_rewrite, protected) - return _process_lifetime_scope(scope, rewrite_broadcast_lifetimes; protected_roots) + checks = Tuple{Symbol,Symbol}[] + guard_roots = setdiff(protected_roots, _assigned_symbols(scope)) + rewritten = InterBroadcastFusion.rewrite_scope(scope; + on_rewrite, protected, guard_roots, alias_checks=checks) + fast = _process_lifetime_scope(rewritten, rewrite_broadcast_lifetimes; protected_roots) + isempty(checks) && return fast + + # Analyze each straight-line branch separately, then wrap the complete + # lifetime-managed bodies. Overlapping inputs retain materialized producers. + fallback = InterBroadcastFusion.rewrite_scope(scope; protected) + slow = _process_lifetime_scope(fallback, rewrite_broadcast_lifetimes; protected_roots) + conditions = [:(cuNumeric._fusion_disjoint($a, $b)) for (a, b) in checks] + condition = reduce((a, b) -> Expr(:&&, a, b), conditions) + return Expr(:if, condition, fast, slow) end diff --git a/src/scoping/inter_broadcast_fusion.jl b/src/scoping/inter_broadcast_fusion.jl index 86730055b..453574ee6 100644 --- a/src/scoping/inter_broadcast_fusion.jl +++ b/src/scoping/inter_broadcast_fusion.jl @@ -15,39 +15,153 @@ using ..ScopingUtils # # The pass is syntax-only and has no NDArray or cuNumeric dependencies. -function _substitute_symbols(expr, replacements::Dict{Symbol,Any}) - assignment = _assignment(expr) - isnothing(assignment) && return _replace_symbols(expr, replacements) - assignment.lhs isa Symbol || return _replace_symbols(expr, replacements) - rhs = _replace_symbols(assignment.rhs, replacements) - return :($(assignment.lhs) = $rhs) +# These records exist only during macro expansion, not during array execution. + +# Array/scalar inputs read by an expression, plus names whose rebinding would +# change its meaning if evaluation is delayed until a later statement. +struct ReadDependencies + inputs::Set{Symbol} + bindings::Set{Symbol} # Also includes callable names that could be rebound. +end + +# A statement such as `tmp = A .+ B` that produces a single-use temporary. +# Records its expression and the statement positions of its definition and use. +struct BroadcastProducer + name::Symbol + definition::Int + consumer::Int + expression::Any +end + +# Temporaries to replace with their expressions at the use site, together with +# any runtime disjointness checks needed to make that delayed evaluation safe. +struct FusionPlan + producers::Dict{Int,BroadcastProducer} + alias_checks::Vector{Tuple{Symbol,Symbol}} +end +FusionPlan() = FusionPlan(Dict{Int,BroadcastProducer}(), Tuple{Symbol,Symbol}[]) + +# Substitutions accumulated while applying a plan, with original statement +# positions and before/after expressions used to report the rewrites. +struct RewriteState + replacements::Dict{Symbol,Any} + sources::Dict{Symbol,Vector{Int}} + events::Vector{NamedTuple} +end +RewriteState() = RewriteState(Dict{Symbol,Any}(), Dict{Symbol,Vector{Int}}(), NamedTuple[]) + +# Only move arithmetic/indexing expressions across other statements. Unknown +# calls may mutate their arguments (even when their names do not end in `!`). +# This name-based allowlist assumes ordinary numerical methods; it does not +# prove that a shadowed function or an overloaded method is free of side effects. +const _READONLY_CALLS = Set((:+, :-, :*, :/, :^, :%, :fld, :cld, :mod, :rem, + :(:), :abs, :abs2, :sqrt, :exp, :log, :sin, :cos, :tan, :inv, :min, :max, + :ifelse, :iszero, :isfinite, :isnan, :identity, :size, :axes, :length, + :firstindex, :lastindex, :eltype, :one, :zero, :<, :>, :<=, :>=, :(==), :(!=), + :Bool, :Int, :Int8, :Int16, :Int32, :Int64, :UInt, :UInt8, :UInt16, :UInt32, + :UInt64, :Float16, :Float32, :Float64, :ComplexF32, :ComplexF64)) + +_read_dependencies!(deps, ::Any) = false +_read_dependencies!(deps, ::Number) = true +_read_dependencies!(deps, ::QuoteNode) = true + +function _read_dependencies!(deps, expr::Symbol) + expr in (:end, :(:), :nothing, :true, :false) || push!(deps, expr) + return true +end + +function _read_dependencies!(deps, expr::Expr) + reference = _reference(expr) + if !isnothing(reference) + return reference.array isa Symbol && + _read_dependencies!(deps, reference.array) && + all(x -> _read_dependencies!(deps, x), reference.indices) + end + # Parameter fields such as args.dt; writes to properties remain barriers. + if expr.head === :. && length(expr.args) == 2 && expr.args[2] isa QuoteNode + return _read_dependencies!(deps, expr.args[1]) + end + call = _call(expr) + isnothing(call) && (call = _dotcall(expr)) + isnothing(call) && return false + f = call.f + _is_broadcast_op(f) && (f = Symbol(chop(string(f); head=1, tail=0))) + return f in _READONLY_CALLS && all(x -> _read_dependencies!(deps, x), call.args) +end + +_is_readonly(expr) = _read_dependencies!(Set{Symbol}(), expr) + +function _dependencies(expr) + inputs = Set{Symbol}() + _read_dependencies!(inputs, expr) || return nothing + return ReadDependencies(inputs, Set(walk_symbols(expr))) end -function _indexed_assignment_base(stmt) +function _statement_assignment(stmt) assignment = _assignment(stmt) - isnothing(assignment) && return nothing - reference = _reference(assignment.lhs) + return isnothing(assignment) ? _broadcast_assignment(stmt) : assignment +end + +# A simple destination whose indexing has no unknown effects. Property writes +# and calls that compute destinations are deliberately excluded. +function _write_target(lhs) + lhs isa Symbol && return lhs + reference = _reference(lhs) isnothing(reference) && return nothing reference.array isa Symbol || return nothing + all(_is_readonly, reference.indices) || return nothing return reference.array end -function _safe_to_delay_broadcast( - stmts, def_idx::Int, use_idx::Int, dependencies::Set{Symbol}, lazy_defs::Set{Int} -) - for i in (def_idx + 1):(use_idx - 1) - stmt = stmts[i] - i in lazy_defs && continue +function _known_assignment(stmt) + assignment = _statement_assignment(stmt) + isnothing(assignment) && return false + return !isnothing(_write_target(assignment.lhs)) && _is_readonly(assignment.rhs) +end + +function _entry_guard_valid(stmts, write_index) + # An earlier unknown call could replace storage after the entry check. + return all(i -> _known_assignment(stmts[i]), 1:(write_index - 1)) +end + +function _write_checks(stmts, index, assignment, deps::ReadDependencies, guard_roots) + _is_readonly(assignment.rhs) || return nothing + target = _write_target(assignment.lhs) + isnothing(target) && return nothing + target in guard_roots || return nothing + _entry_guard_valid(stmts, index) || return nothing + # Entry checks can name stable arguments, not locals created later. + issubset(deps.inputs, guard_roots) || return nothing + target in deps.inputs && return nothing + return [(target, input) for input in Base.sort!(collect(deps.inputs); by=string)] +end - # An indexed write to an unrelated array does not invalidate the lazy - # producer. Any other intervening statement is conservatively a barrier. - mutated = _indexed_assignment_base(stmt) - if !isnothing(mutated) && !(mutated in dependencies) +function _delay_checks(stmts, producer::BroadcastProducer, rhs, guard_roots) + deps = _dependencies(rhs) + isnothing(deps) && return nothing + checks = Tuple{Symbol,Symbol}[] + for i in (producer.definition + 1):(producer.consumer - 1) + assignment = _assignment(stmts[i]) + if !isnothing(assignment) && assignment.lhs isa Symbol + assignment.lhs in deps.bindings && return nothing + _is_readonly(assignment.rhs) || return nothing continue end - return false + isnothing(assignment) && (assignment = _broadcast_assignment(stmts[i])) + isnothing(assignment) && return nothing + write_checks = _write_checks(stmts, i, assignment, deps, guard_roots) + isnothing(write_checks) && return nothing + append!(checks, write_checks) end - return true + return checks +end + +function _substitute_symbols(expr, replacements::Dict{Symbol,Any}) + assignment = _assignment(expr) + isnothing(assignment) && return _replace_symbols(expr, replacements) + assignment.lhs isa Symbol || return _replace_symbols(expr, replacements) + rhs = _replace_symbols(assignment.rhs, replacements) + return :($(assignment.lhs) = $rhs) end function _single_use_index(stmts, symbol::Symbol, def_idx::Int) @@ -89,68 +203,85 @@ function _fuse_into_destination(stmt) return Expr(:(.=), assignment.lhs, assignment.rhs) end -function _rewrite_scope(scope, protected) - stmts = _scope_statements(scope) - isnothing(stmts) && return scope, NamedTuple[] - - definitions = Dict{Symbol,Tuple{Int,Any}}() - lazy_defs = Set{Int}() - for (i, stmt) in enumerate(stmts) - assignment = _assignment(stmt) - if !isnothing(assignment) && assignment.lhs isa Symbol && - _is_broadcast_syntax(assignment.rhs) - definitions[assignment.lhs] = (i, assignment.rhs) - push!(lazy_defs, i) - end - end +function _broadcast_consumer(stmt) + assignment = _statement_assignment(stmt) + rhs = isnothing(assignment) ? stmt : assignment.rhs + _is_broadcast_syntax(rhs) && _is_readonly(rhs) || return false + return isnothing(assignment) || !isnothing(_write_target(assignment.lhs)) +end - inlineable = Dict{Symbol,Tuple{Int,Any}}() - for (sym, (def_idx, rhs)) in definitions - # Never inline a returned producer; it must escape as a real NDArray. - sym in protected && continue - use_idx = _single_use_index(stmts, sym, def_idx) - isnothing(use_idx) && continue - dependencies = Set(walk_symbols(rhs)) - if !_safe_to_delay_broadcast(stmts, def_idx, use_idx, dependencies, lazy_defs) - continue - end - inlineable[sym] = (def_idx, rhs) +function _producer(stmts, index, protected) + assignment = _assignment(stmts[index]) + isnothing(assignment) && return nothing + name, rhs = assignment.lhs, assignment.rhs + name isa Symbol && _is_broadcast_syntax(rhs) || return nothing + name in protected && return nothing + consumer = _single_use_index(stmts, name, index) + isnothing(consumer) && return nothing + _broadcast_consumer(stmts[consumer]) || return nothing + return BroadcastProducer(name, index, consumer, rhs) +end + +function _plan_fusion(stmts, protected, guard_roots) + plan = FusionPlan() + expanded = Dict{Symbol,Any}() + for index in eachindex(stmts) + producer = _producer(stmts, index, protected) + isnothing(producer) && continue + # Include the original inputs of already-elided producers. Otherwise a + # later rebind can become invisible through a chain like t -> s -> out. + rhs = _replace_symbols(producer.expression, expanded) + checks = _delay_checks(stmts, producer, rhs, guard_roots) + isnothing(checks) && continue + plan.producers[index] = producer + expanded[producer.name] = rhs + append!(plan.alias_checks, checks) end + unique!(plan.alias_checks) + return plan +end - replacements = Dict{Symbol,Any}() - replacement_sources = Dict{Symbol,Vector{Int}}() - removed = Set(first(info) for info in values(inlineable)) - def_symbols = Dict(info[1] => sym for (sym, info) in inlineable) - fusion_events = NamedTuple[] - rewritten = Any[] +function _record_producer!(state::RewriteState, producer::BroadcastProducer) + indices = _source_indices(producer.expression, state.sources) + push!(indices, producer.definition) + state.sources[producer.name] = indices + state.replacements[producer.name] = + _substitute_symbols(producer.expression, state.replacements) + return nothing +end - for (i, original_stmt) in enumerate(stmts) - if i in removed - sym = def_symbols[i] - assignment = _assignment(original_stmt) - source_indices = _source_indices(assignment.rhs, replacement_sources) - push!(source_indices, i) - replacement_sources[sym] = source_indices - replacements[sym] = _substitute_symbols(inlineable[sym][2], replacements) - continue - end +function _rewrite_consumer!(state::RewriteState, stmts, index) + original = stmts[index] + indices = _source_indices(original, state.sources) + stmt = _substitute_symbols(original, state.replacements) + isempty(indices) && return stmt + + stmt = _fuse_into_destination(stmt) + before = Expr(:block, (stmts[i] for i in indices)..., original) + push!(state.events, (; before, fused=stmt)) + return stmt +end - source_indices = _source_indices(original_stmt, replacement_sources) - stmt = _substitute_symbols(original_stmt, replacements) - - if !isempty(source_indices) - stmt = _fuse_into_destination(stmt) - before = Expr( - :block, - (stmts[source_idx] for source_idx in source_indices)..., - original_stmt, - ) - push!(fusion_events, (; before, fused=stmt)) +function _apply_plan(scope, stmts, plan::FusionPlan) + state = RewriteState() + rewritten = Any[] + for index in eachindex(stmts) + producer = get(plan.producers, index, nothing) + if isnothing(producer) + push!(rewritten, _rewrite_consumer!(state, stmts, index)) + else + _record_producer!(state, producer) end - push!(rewritten, stmt) end + return Expr(scope.head, rewritten...), state.events +end - return Expr(scope.head, rewritten...), fusion_events +function _rewrite_scope(scope, protected, guard_roots) + stmts = _scope_statements(scope) + isnothing(stmts) && return scope, NamedTuple[], Tuple{Symbol,Symbol}[] + plan = _plan_fusion(stmts, protected, guard_roots) + rewritten, events = _apply_plan(scope, stmts, plan) + return rewritten, events, plan.alias_checks end """ @@ -161,9 +292,18 @@ rewritten scope. Symbols in `protected` — typically whatever the scope returns are never fused so they stay materialized. When provided, `on_rewrite` is called with a named tuple containing the `before` and `fused` expressions for each rewrite. + +Nonadjacent producers may cross read-only bindings that do not rebind their +inputs. With `guard_roots` and an `alias_checks` collector, writes to stable +arguments may also be crossed: the caller must guard the returned rewrite with +disjointness checks for those `(destination, input)` pairs and provide a fallback. """ -function rewrite_scope(scope; on_rewrite=nothing, protected=Set{Symbol}()) - rewritten, fusion_events = _rewrite_scope(scope, protected) +function rewrite_scope(scope; on_rewrite=nothing, protected=Set{Symbol}(), + guard_roots=Set{Symbol}(), alias_checks=nothing) + # Callers without a guard collector receive only statically safe rewrites. + roots = isnothing(alias_checks) ? Set{Symbol}() : guard_roots + rewritten, fusion_events, checks = _rewrite_scope(scope, protected, roots) + isnothing(alias_checks) || append!(alias_checks, checks) if !isnothing(on_rewrite) for event in fusion_events on_rewrite(event) diff --git a/src/scoping/lifetimes.jl b/src/scoping/lifetimes.jl index b13bf5d35..1a82aee5e 100644 --- a/src/scoping/lifetimes.jl +++ b/src/scoping/lifetimes.jl @@ -39,7 +39,13 @@ function rewrite_eager_lifetimes(scope) if !isnothing(broadcast_assignment) (; lhs, rhs) = broadcast_assignment op = expr.head - new_lhs, lhs_temps = rewrite(lhs) + # Indexed broadcast assignment needs a view even when getindex + # copies (notably all-colon indexing). + new_lhs, lhs_temps = if isnothing(_reference(lhs)) + rewrite(lhs) + else + fresh_tmp(:(Base.@view $lhs)) + end # Do not hoist the top-level call of the RHS to preserve fusion. call = _call(rhs) if !isnothing(call) diff --git a/test/analysis/accelerate.jl b/test/analysis/accelerate.jl index 41f8d087f..2c0038895 100644 --- a/test/analysis/accelerate.jl +++ b/test/analysis/accelerate.jl @@ -25,6 +25,33 @@ using InteractiveUtils: code_typed +@testset "@accelerate respects rebindings and alias writes" begin + @accelerate function _acc_rebind(a, b) + t = a .+ 1f0 + a = b .+ 2f0 + t .+ a + end + @accelerate function _acc_aliaswrite(a, b) + t = a .+ 1f0 + b[1] = 9f0 + t .+ 0f0 + end + for make in (identity, NDArray) + @test Array(_acc_rebind(make(Float32[1]), make(Float32[10]))) == Float32[14] + a = make(Float32[1, 2]) + @allowscalar result = _acc_aliaswrite(a, a) + @test Array(result) == Float32[2, 3] + end + # Adjacent chains still inline their single-use intermediates. + ex = cuNumeric.InterBroadcastFusion.rewrite_scope(quote + t = a .+ 1f0 + u = t .* 2f0 + u .+ 3f0 + end) + @test !(:t in cuNumeric.ScopingUtils.walk_symbols(ex)) + @test !(:u in cuNumeric.ScopingUtils.walk_symbols(ex)) +end + @testset "@accelerate — four forms" begin T = Float32 N = 64 diff --git a/test/analysis/inter_broadcast_fusion.jl b/test/analysis/inter_broadcast_fusion.jl new file mode 100644 index 000000000..aebe96c94 --- /dev/null +++ b/test/analysis/inter_broadcast_fusion.jl @@ -0,0 +1,143 @@ +using Test + +const IBF = cuNumeric.InterBroadcastFusion +const SU = cuNumeric.ScopingUtils + +@testset "Nonadjacent producer elimination" begin + ex = IBF.rewrite_scope(quote + t = a .+ 1f0 + r = b .* 2f0 + out .= t .+ r + end) + @test !(:t in SU.walk_symbols(ex)) + @test !(:r in SU.walk_symbols(ex)) + + # Expanding t into s must retain a as a dependency of s. + ex = IBF.rewrite_scope(quote + t = a .+ 1f0 + s = t .* 2f0 + a = b .+ 2f0 + s .+ a + end) + @test !(:t in SU.walk_symbols(ex)) + @test :s in SU.walk_symbols(ex) + + # Calls can mutate inputs without a bang suffix, including inside a broadcast. + for barrier in (:(touch(a)), :(unused = touch.(a)), :(a[1] = 9f0)) + ex = IBF.rewrite_scope(quote + t = a .+ 1f0 + $barrier + out .= t .* 2f0 + end) + @test :t in SU.walk_symbols(ex) + end + ex = IBF.rewrite_scope(quote + t = a .+ 1f0 + r = b .* 2f0 + t .+ r + end; protected=Set([:t])) + @test :t in SU.walk_symbols(ex) + + ex = IBF.rewrite_scope(quote + t = a .+ 1f0 + out[touch(a)] .= t .* 2f0 + end) + @test :t in SU.walk_symbols(ex) + ex = IBF.rewrite_scope(quote + t = sin.(a) + sin = cos + t .+ 1f0 + end) + @test :t in SU.walk_symbols(ex) + + # A compact Gray-Scott-shaped sequence: four producers and two array writes. + body = quote + F_u = u .* v + F_v = u .+ v + u_lap = u .* 2f0 + v_lap = v .* 3f0 + u_new[:] = F_u .+ u_lap + v_new[:] = F_v .+ v_lap + nothing + end + checks = Tuple{Symbol,Symbol}[] + ex = IBF.rewrite_scope(body; guard_roots=Set([:u, :v, :u_new, :v_new]), + alias_checks=checks) + @test all(s -> !(s in SU.walk_symbols(ex)), (:F_u, :F_v, :u_lap, :v_lap)) + @test Set(checks) == Set([(:u_new, :u), (:u_new, :v)]) + fallback = IBF.rewrite_scope(body) + @test :F_v in SU.walk_symbols(fallback) + @test :v_lap in SU.walk_symbols(fallback) + + # Guards cannot be hoisted past unknown calls that could change storage. + prefixed = Expr(:block, :(prepare(u_new, u)), SU._scope_statements(body)...) + checks = Tuple{Symbol,Symbol}[] + IBF.rewrite_scope(prefixed; guard_roots=Set([:u, :v, :u_new, :v_new]), alias_checks=checks) + @test isempty(checks) +end + +@accelerate function _acc_two_updates!(u, v, u_new, v_new) + F_u = u .* v + F_v = u .+ v + u_lap = u .* 2f0 + v_lap = v .* 3f0 + u_new[:] = F_u .+ u_lap + v_new[:] = F_v .+ v_lap + nothing +end +function _plain_two_updates!(u, v, u_new, v_new) + F_u = u .* v + F_v = u .+ v + u_lap = u .* 2f0 + v_lap = v .* 3f0 + u_new[:] = F_u .+ u_lap + v_new[:] = F_v .+ v_lap + nothing +end + +@testset "Guarded fusion preserves shared inputs" begin + for make in (identity, NDArray), alias in (:none, :u, :v, :shifted) + host = Float32[1, 2, 3, 4, 5] + hu = view(host, 1:4) + hv = Float32[5, 6, 7, 8] + ho = alias === :none ? zeros(Float32, 4) : + alias === :u ? hu : alias === :v ? hv : view(host, 2:5) + hout = zeros(Float32, 4) + parent = make(copy(host)) + u = view(parent, 1:4) + v = make(copy(hv)) + unew = alias === :none ? make(zeros(Float32, 4)) : + alias === :u ? u : alias === :v ? v : view(parent, 2:5) + vnew = make(zeros(Float32, 4)) + _plain_two_updates!(hu, hv, ho, hout) + @test _acc_two_updates!(u, v, unew, vnew) === nothing + @test Array(unew) == ho + @test Array(vnew) == hout + @test Array(parent) == host + if make === NDArray + foreach(cuNumeric.destroy!, (unew, vnew, u, v, parent)) + end + end +end + +@testset "Transitive rebindings and unknown calls remain barriers" begin + @accelerate function transitive(a, b) + t = a .+ 1f0 + s = t .* 2f0 + a = b .+ 2f0 + s .+ a + end + function touch(a) + fill!(a, 9f0) + nothing + end + @accelerate function unknown_effect(a) + t = a .+ 1f0 + touch(a) + t .* 2f0 + end + for make in (identity, NDArray) + @test Array(transitive(make(Float32[1, 2]), make(Float32[10, 20]))) == Float32[16, 28] + @test Array(unknown_effect(make(Float32[1, 2]))) == Float32[4, 6] + end +end diff --git a/test/array/conversion_lifetimes.jl b/test/array/conversion_lifetimes.jl index 72ba9e462..34b0fb65f 100644 --- a/test/array/conversion_lifetimes.jl +++ b/test/array/conversion_lifetimes.jl @@ -1,5 +1,23 @@ using Test +@testset "Direct host-vector copy ownership" begin + for T in (Float32, Float64, ComplexF32, Int32, Bool), n in (0, 1, 4) + source = fill(one(T), n) + dest = cuNumeric.zeros(T, n) + @test copyto!(dest, source) === dest + fill!(source, zero(T)) + GC.gc(true) + @test Array(dest) == fill(one(T), n) + @test_throws DimensionMismatch copyto!(dest, fill(one(T), n + 1)) + end + parent = cuNumeric.zeros(Float32, 6) + dest = view(parent, 2:5) + source = Float32[1, 2, 3, 4] + copyto!(dest, source) + fill!(source, 9f0) + @test Array(parent) == Float32[0, 1, 2, 3, 4, 0] +end + @testset "copyto! from Array" begin expected = reshape(ComplexF64.(1:8), 2, 2, 2) source = copy(expected) diff --git a/test/array/diagonal_updates.jl b/test/array/diagonal_updates.jl new file mode 100644 index 000000000..6e4999914 --- /dev/null +++ b/test/array/diagonal_updates.jl @@ -0,0 +1,167 @@ +using Test, LinearAlgebra + +@testset "Diagonal division writes existing destination storage" begin + @allowpromotion for T in (Float32, Float64, ComplexF32, ComplexF64) + dh = T <: Complex ? T[2 + im, 3 - im, 4 + 2im] : T[2, 3, 4] + D = Diagonal(NDArray(dh)) + for ah in (T[2, 6, 12], reshape(T.(1:6), 3, 2)) + a = NDArray(ah) + alias = view(a, ntuple(_ -> Colon(), ndims(a))...) + expected = ah ./ (ndims(ah) == 1 ? dh : reshape(dh, 3, 1)) + @test Array(D \ a) ≈ expected + @test ldiv!(D, a) === a + @test Array(alias) ≈ expected + end + ah = reshape(T.(1:6), 2, 3) + a = NDArray(ah) + alias = view(a, :, :) + expected = ah ./ reshape(dh, 1, 3) + @test Array(a / D) ≈ expected + @test rdiv!(a, D) === a + @test Array(alias) ≈ expected + @test Array(D.diag) == dh + end + + # A shared diagonal must be read completely before its storage is overwritten. + for divide! in (ldiv!, rdiv!) + host = reshape(Float32.(1:9), 3, 3) + a = NDArray(host) + # NDArray slicing retains singleton dimensions; Diagonal needs a vector. + D = Diagonal(cuNumeric.reshape(view(a, :, 1), (3,))) + @test cuNumeric.nda_overlaps(a, D.diag) + dh = copy(host[:, 1]) + expected = host ./ reshape(dh, divide! === ldiv! ? (3, 1) : (1, 3)) + divide! === ldiv! ? ldiv!(D, a) : rdiv!(a, D) + @test Array(a) ≈ expected + @test Array(D.diag) ≈ expected[:, 1] + end + parent = NDArray(Float32[2, 4, 8, 16]) + @test ldiv!(Diagonal(view(parent, 1:3)), view(parent, 2:4)) isa NDArray + @test Array(parent) == Float32[2, 2, 2, 2] + + @allowpromotion for (AType, DType) in ((Float32, Float64), (Float64, Float32), (Int32, Float32)) + dh = DType[2, 4] + ah = AType[4 8; 8 16] + D = Diagonal(NDArray(dh)) + a, b = NDArray(ah), NDArray(ah) + @test ldiv!(D, a) === a + @test rdiv!(b, D) === b + @test Array(a) == AType.(ah ./ reshape(dh, 2, 1)) + @test Array(b) == AType.(ah ./ reshape(dh, 1, 2)) + end + + for T in (Float32, Float64) + dh = T[0, -0.0, Inf, NaN] + ah = reshape(T[1, 0, -1, Inf, 0, 1, Inf, NaN], 2, 4) + D = Diagonal(NDArray(dh)) + a = NDArray(ah) + expected = ah ./ reshape(dh, 1, 4) + @test isequal(Array(a / D), expected) + @test rdiv!(a, D) === a + @test isequal(Array(a), expected) + b = NDArray(copy(permutedims(ah))) + @test ldiv!(D, b) === b + @test isequal(Array(b), permutedims(expected)) + end + + D = Diagonal(NDArray(Float32[2])) + @test_throws DimensionMismatch ldiv!(D, cuNumeric.ones(Float32, 3)) + @test_throws DimensionMismatch ldiv!(D, cuNumeric.ones(Float32, 3, 2)) + @test_throws DimensionMismatch rdiv!(cuNumeric.ones(Float32, 2, 3), D) + empty = Diagonal(NDArray(Float32[])) + @test size(ldiv!(empty, cuNumeric.zeros(Float32, 0, 2))) == (0, 2) + @test size(rdiv!(cuNumeric.zeros(Float32, 2, 0), empty)) == (2, 0) +end + +@testset "Diagonal products write existing destination storage" begin + @allowpromotion for T in (Float32, Float64, ComplexF32, ComplexF64) + dh = T[2, 3, 4] + D = Diagonal(NDArray(dh)) + for ah in (T[1, 2, 3], reshape(T.(1:6), 3, 2)) + a = NDArray(ah) + c = cuNumeric.zeros(T, size(ah)) + v = view(c, ntuple(_ -> Colon(), ndims(c))...) + @test mul!(c, D, a) === c + @test Array(v) ≈ Diagonal(dh) * ah + @test lmul!(D, a) === a + @test Array(a) ≈ Diagonal(dh) * ah + end + ah = reshape(T.(1:6), 2, 3) + a = NDArray(ah) + c = cuNumeric.zeros(T, 2, 3) + v = view(c, :, :) + @test mul!(c, a, D) === c + @test Array(v) ≈ ah * Diagonal(dh) + @test rmul!(a, D) === a + @test Array(a) ≈ ah * Diagonal(dh) + @test_throws DimensionMismatch mul!(cuNumeric.zeros(T, 2), D, NDArray(T[1, 2, 3])) + @test_throws DimensionMismatch mul!(cuNumeric.zeros(T, 2, 2), cuNumeric.ones(T, 2, 2), D) + end + + # Partially overlapping vector input and output must use a temporary. + a = NDArray(Float32[1, 2, 3, 4]) + D = Diagonal(NDArray(Float32[2, 3, 4])) + mul!(view(a, 2:4), D, view(a, 1:3)) + @test Array(a) == Float32[1, 2, 6, 12] + + @allowpromotion begin + c = cuNumeric.zeros(Float64, 3) + mul!(c, D, NDArray(Float32[1, 2, 3])) + @test Array(c) == [2.0, 6.0, 12.0] + end +end + +@testset "Diagonal norm dispatch and autofetch policy" begin + dh = Float32[3, 4] + D = Diagonal(NDArray(dh)) + for p in (0, 0f0, -0.0, 1, 1f0, 2, 2f0, 2.0, 3, 1.5, Inf, Inf32, -Inf32) + actual = fetch(norm(D, p)) + expected = norm(Diagonal(dh), p) + @test actual ≈ expected + @test typeof(actual) == typeof(expected) + end + for p in (NDArray(2), cnscalar(NDArray(2))) + @test_throws "Implicit CNScalar host extraction is disabled" norm(D, p) + @allowautofetch @test fetch(norm(D, p)) ≈ 5f0 + @test fetch(norm(D, fetch(p))) ≈ 5f0 + end +end + +@testset "Diagonal condition results survive temporary cleanup" begin + @allowpromotion for T in (Float32, Float64, ComplexF32, ComplexF64) + dh = T <: Complex ? T[1 + im, 2 - im, 4 + im] : T[1, 2, 4] + D = Diagonal(NDArray(dh)) + orders = (1, 2, Inf) + results = map(p -> cond(D, p), orders) + # Fetch only after temporary handles have been released and collected. + GC.gc() + cuNumeric.drain_pending_frees!() + for (p, result) in zip(orders, results) + @test fetch(result) ≈ cond(Diagonal(dh), p) + end + @test Array(D.diag) == dh + end +end + +@testset "Diagonal zero and negative norms" begin + @allowpromotion for T in (Float32, Float64, ComplexF32, ComplexF64) + for dh in (T[], T[2], T[0], T[2, 3], T[2, 0], T[NaN, 2], T[Inf, 2]) + D = Diagonal(NDArray(dh)) + for p in (0, -1, -2, -Inf) + expected = norm(Diagonal(dh), p) + actual = fetch(norm(D, p)) + @test isapprox(actual, expected; nans=true) + @test typeof(actual) == typeof(expected) + end + end + end +end + +@testset "Diagonal product includes structural zeros" begin + @allowpromotion for T in (Float32, Float64, ComplexF32, ComplexF64) + for dh in (T[], T[2], T[0], T[2, 3], T[NaN, 2], T[Inf, 2], + T[floatmax(real(T)), floatmax(real(T))]) + @test isequal(fetch(prod(Diagonal(NDArray(dh)))), prod(Diagonal(dh))) + end + end +end diff --git a/test/array/sort.jl b/test/array/sort.jl index 7a61a227b..18ead1138 100644 --- a/test/array/sort.jl +++ b/test/array/sort.jl @@ -132,6 +132,18 @@ end end end +@testset "Scalar search preserves query precision" begin + for (h, queries) in ((Int64[1, 3, 5], (2.5, -0.5, 5.5)), + (Float32[1, 2, 3], (1.0 + eps(Float64), 2.0 - eps(Float64), 4.0))) + a = NDArray(h) + for q in queries + @test fetch(cuNumeric.searchsortedfirst(a, q)) == Base.searchsortedfirst(h, q) + @test fetch(cuNumeric.searchsortedlast(a, q)) == Base.searchsortedlast(h, q) + @test cuNumeric.searchsorted(a, q) == Base.searchsorted(h, q) + end + end +end + @testset "unique" begin @testset verbose = true for T in SORT_TYPES A = _unique_fixture(T) diff --git a/test/array/storage_semantics.jl b/test/array/storage_semantics.jl new file mode 100644 index 000000000..50443222e --- /dev/null +++ b/test/array/storage_semantics.jl @@ -0,0 +1,95 @@ +using Test + +@testset "Shape queries across ranks and shared handles" begin + for dims in ((), (0,), (5,), (2, 3), (2, 1, 3)) + a = cuNumeric.zeros(Float32, dims) + @test size(a) === dims + @test axes(a) == map(Base.OneTo, dims) + @test length(a) == prod(dims) + cuNumeric.destroy!(a) + end + a = cuNumeric.zeros(Float32, 4, 5) + sliced = view(a, 2:3, 1:4) + reshaped = cuNumeric.reshape(a, (2, 2, 5)) + @test size(sliced) === (2, 4) + @test size(reshaped) === (2, 2, 5) + @test size(a) === (4, 5) + foreach(cuNumeric.destroy!, (sliced, reshaped, a)) +end + +@testset "Full-colon views share storage and own their handles" begin + for shape in ((4,), (2, 3), (2, 2, 3)) + a = cuNumeric.zeros(Float32, shape) + inds = ntuple(_ -> Colon(), length(shape)) + v = view(a, inds...) + @test v !== a + @test size(v) == shape + fill!(v, 3f0) + @test Array(a) == fill(3f0, shape) + fill!(a, 4f0) + @test Array(v) == fill(4f0, shape) + cuNumeric.destroy!(v) + a[inds...] .= 5f0 + @test Array(a) == fill(5f0, shape) + copied = a[inds...] + fill!(copied, 9f0) + @test Array(a) == fill(5f0, shape) + end + a = cuNumeric.zeros(Float32, 0) + @test size(view(a, :)) == (0,) + + a = NDArray(2f0) + v = view(a) + @test v !== a + @test size(v) == () + fill!(v, 3f0) + @test fetch(a) == 3f0 + cuNumeric.destroy!(v) + @test fetch(a) == 3f0 + + @accelerate function _full_view_update!(a) + a[:, :] .= 7f0 + a + end + a = cuNumeric.zeros(Float32, 2, 3) + @test _full_view_update!(a) === a + @test Array(a) == fill(7f0, 2, 3) +end + +@testset "In-place broadcast preserves existing views" begin + a = NDArray(Float32[1, 2, 3, 4]) + v = view(a, 1:2) + a .+= 1f0 + @test Array(v) == Float32[2, 3] + a .= a .* 2f0 .+ 1f0 + @test Array(v) == Float32[5, 7] + + reshaped = cuNumeric.reshape(a, (2, 2)) + reshaped .+= 1f0 + @test Array(a) == Float32[6, 8, 10, 12] + @test Array(v) == Float32[6, 8] + + # Changing result type must copy back into the existing destination store. + ints = NDArray(Int32[1, 2, 3, 4]) + a .= ints .+ Int32(1) + @test Array(v) == Float32[2, 3] + + # Shifted inputs need a temporary, including identity broadcasts. + a = NDArray(Float32[1, 2, 3, 4]) + a[2:4] .= a[1:3] + @test Array(a) == Float32[1, 1, 2, 3] + a[2:4] .= a[1:3] .+ 10f0 + @test Array(a) == Float32[1, 11, 11, 12] +end + +@testset "Reject unsupported multidimensional linear indexing" begin + a = NDArray(Float32[1 3 5; 2 4 6]) + @test_throws ArgumentError a[1:2] + @test_throws ArgumentError a[:] + @test_throws ArgumentError view(a, :) + @test_throws ArgumentError (a[1:2] = NDArray(Float32[9, 9])) + @test_throws ArgumentError (a[:] = cuNumeric.ones(Float32, 2, 3)) + @test_throws ArgumentError (a[:] .= 9f0) + @test Array(a[1:2, 2:3]) == Float32[3 5; 4 6] + @test Array(a[:, 2]) == reshape(Float32[3, 4], 2, 1) +end diff --git a/test/array/unfused_destinations.jl b/test/array/unfused_destinations.jl new file mode 100644 index 000000000..acdbf434d --- /dev/null +++ b/test/array/unfused_destinations.jl @@ -0,0 +1,67 @@ +using Test + +@testset "Native n-ary broadcasts reuse the destination" begin + a = NDArray(Float32[1, 2, 3, 4]) + b = NDArray(Float32[2, 3, 4, 5]) + c = NDArray(Float32[3, 4, 5, 6]) + dest = cuNumeric.zeros(Float32, 4) + for f in (+, *) + for args in ((a, b, c), (a, b, c, a), (2f0, 3f0, a), + (Base.broadcasted(+, a, 1f0), b, c)) + bc = Base.Broadcast.instantiate(Base.broadcasted(f, args...)) + host_args = map(args) do arg + arg isa NDArray && return Array(arg) + arg isa Base.Broadcast.Broadcasted && return Array(a) .+ 1f0 + return arg + end + expected = f.(host_args...) + # Exercise the native route regardless of the fusion preference. + result = cuNumeric.unravel_broadcast_tree(bc, dest) + @test result === dest + @test Array(dest) == expected + @test Array(a) == Float32[1, 2, 3, 4] + @test Array(b) == Float32[2, 3, 4, 5] + @test Array(c) == Float32[3, 4, 5, 6] + + allocated = cuNumeric.unravel_broadcast_tree(bc) + @test allocated !== dest + @test Array(allocated) == expected + cuNumeric.destroy!(allocated) + end + # The destination can be read both before and during the final operation. + fill!(dest, 2f0) + bc = Base.Broadcast.instantiate(Base.broadcasted(f, dest, a, dest)) + @test cuNumeric._unfused_into!(dest, bc) === dest + @test Array(dest) == f.(2f0, Float32[1, 2, 3, 4], 2f0) + end + foreach(cuNumeric.destroy!, (a, b, c, dest)) +end + +@testset "Native n-ary broadcasts preserve overlap and conversion guards" begin + for f in (+, *) + parent = NDArray(Float32[1, 2, 3, 4, 5]) + source = view(parent, 1:4) + dest = view(parent, 2:5) + a = cuNumeric.fill(2f0, 4) + b = cuNumeric.fill(3f0, 4) + bc = Base.Broadcast.instantiate(Base.broadcasted(f, a, b, source)) + result = cuNumeric.unravel_broadcast_tree(bc, dest) + @test result !== dest + @test Array(parent) == Float32[1, 2, 3, 4, 5] + @test cuNumeric._copyto_unfused!(dest, result) === dest + @test Array(parent) == vcat(1f0, f.(2f0, 3f0, Float32[1, 2, 3, 4])) + foreach(cuNumeric.destroy!, (source, dest, parent, a, b)) + end + + ints = NDArray(Int32[1, 2, 3, 4]) + dest = cuNumeric.zeros(Float32, 4) + alias = view(dest, 1:2) + bc = Base.Broadcast.instantiate(Base.broadcasted(+, ints, Int32(2), Int32(3))) + result = cuNumeric.unravel_broadcast_tree(bc, dest) + @test result !== dest + @test eltype(result) === Int32 + @test cuNumeric._copyto_unfused!(dest, result) === dest + @test Array(dest) == Float32[6, 7, 8, 9] + @test Array(alias) == Float32[6, 7] + foreach(cuNumeric.destroy!, (alias, dest, ints)) +end