Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
89 changes: 52 additions & 37 deletions src/execution.jl
Original file line number Diff line number Diff line change
Expand Up @@ -114,16 +114,7 @@ function _success_gate(prepared::PreparedPlan, lease_index::Int32, parent)
expected = prepared.owner,
actual = current_task(),
))
statuses = map(
status -> status.device, _prepared_validation_statuses(prepared)
)
isempty(statuses) && throw(LocalMathValidationError(
"success_gate requires a source law with device validation status";
stage = :prepare,
contract = :validation_status,
expected = :device_validation_status,
actual = :none,
))
statuses = (prepared.runtime.execution_gate,)
return _SuccessfulLawGate(
parent, statuses, lease_index
)
Expand Down Expand Up @@ -411,43 +402,59 @@ function _cache_receipt_result!(receipt::ExecutionReceipt)
return failure
end

function _synchronize_receipt_scope!(receipt::ExecutionReceipt)
lane = receipt.prepared.lane
receipt.scope_ordinal <= _lane_settled_ordinal(lane) && return nothing
target = _lane_scope_ordinal(lane)
try
_wait_lane!(lane)
_mark_lane_settled!(lane, target)
catch error
annotated = _provider_execution_error(error, :wait)
receipt.state = _EXECUTION_RECEIPT_PROVIDER_FAILURE
receipt.failure = annotated
throw(annotated)
end
function _settle_prepared_scope!(prepared::PreparedPlan, target::UInt64,
seen::Base.IdSet{Any})
lane = prepared.lane
target <= _lane_settled_ordinal(lane) && return nothing
runtime = prepared.runtime
_settle_lane_tail!(lane, runtime.execution_gate, runtime.validation_host)
push!(seen, prepared)
_mark_lane_settled!(lane, target)
return nothing
end

function _transfer_receipt_statuses!(receipt::ExecutionReceipt, seen::Base.IdSet{Any})
prepared = receipt.prepared
if !(prepared in seen)
_transfer_settled_validation_statuses!(prepared.lane,
_prepared_validation_statuses(prepared))
runtime = prepared.runtime
_transfer_settled_validation_status!(prepared.lane,
runtime.execution_gate, runtime.validation_host)
push!(seen, prepared)
end
for dependency in receipt.dependencies
# Admission permits unresolved dependencies only within this provider
# scope; settled cross-scope dependencies already own host-visible status.
_receipt_settled(dependency) ||
_transfer_receipt_statuses!(dependency, seen)
end
return nothing
end

function _settle_receipt_statuses!(receipt::ExecutionReceipt,
seen::Base.IdSet{Any})
try
lane = receipt.prepared.lane
receipt.scope_ordinal <= _lane_settled_ordinal(lane) ||
_settle_prepared_scope!(receipt.prepared,
_lane_scope_ordinal(lane), seen)
_transfer_receipt_statuses!(receipt, seen)
catch error
annotated = _provider_execution_error(error, :wait)
receipt.state = _EXECUTION_RECEIPT_PROVIDER_FAILURE
receipt.failure = annotated
_release_receipt_lease!(receipt)
throw(annotated)
end
return nothing
end

function Base.wait(receipt::ExecutionReceipt)
current_task() === receipt.prepared.owner || throw(LocalMathValidationError(
"this ExecutionReceipt belongs to another owner task";
stage = :wait, contract = :receipt_owner))
if !_receipt_settled(receipt)
_synchronize_receipt_scope!(receipt)
_transfer_receipt_statuses!(receipt, Base.IdSet{Any}())
seen = Base.IdSet{Any}()
_settle_receipt_statuses!(receipt, seen)
_cache_receipt_result!(receipt)
end
receipt.failure === nothing || throw(receipt.failure)
Expand All @@ -457,16 +464,16 @@ end
"""
waitall(receipts::ExecutionReceipt...)

Settle several logical receipts, synchronizing each represented provider scope
at most once. Cached failures are reported deterministically in argument order.
Settle several logical receipts, completing each represented provider scope at
most once. Cached failures are reported deterministically in argument order.
Inspection is not required before waiting and waiting remains idempotent.
"""
function waitall(receipts::Tuple{Vararg{ExecutionReceipt}})
isempty(receipts) && throw(LocalMathValidationError(
"waitall requires at least one ExecutionReceipt";
stage = :wait, contract = :receipt_group_arity,
expected = :nonempty, actual = 0))
lanes = Any[]
settlement_prepared = Any[]
targets = UInt64[]
for receipt in receipts
current_task() === receipt.prepared.owner || throw(
Expand All @@ -475,29 +482,36 @@ function waitall(receipts::Tuple{Vararg{ExecutionReceipt}})
stage = :wait, contract = :receipt_owner))
lane = receipt.prepared.lane
index = findfirst(candidate ->
_lane_same_wait_scope(candidate, lane), lanes)
_lane_same_wait_scope(candidate.lane, lane), settlement_prepared)
if index === nothing
push!(lanes, lane)
push!(settlement_prepared, receipt.prepared)
push!(targets, _lane_scope_ordinal(lane))
else
targets[index] = max(targets[index], _lane_scope_ordinal(lane))
end
end
provider_failures = IdDict{Any,Any}()
for (lane, target) in zip(lanes, targets)
seen = Base.IdSet{Any}()
for (prepared, target) in zip(settlement_prepared, targets)
lane = prepared.lane
target <= _lane_settled_ordinal(lane) && continue
try
_wait_lane!(lane)
_mark_lane_settled!(lane, target)
_settle_prepared_scope!(prepared, target, seen)
catch error
provider_failures[lane.scope] =
_provider_execution_error(error, :wait)
end
end
seen = Base.IdSet{Any}()
for receipt in receipts
haskey(provider_failures, receipt.prepared.lane.scope) && continue
_receipt_settled(receipt) || _transfer_receipt_statuses!(receipt, seen)
if !_receipt_settled(receipt)
try
_transfer_receipt_statuses!(receipt, seen)
catch error
provider_failures[receipt.prepared.lane.scope] =
_provider_execution_error(error, :wait)
end
end
end
first_failure = nothing
for receipt in receipts
Expand All @@ -509,6 +523,7 @@ function waitall(receipts::Tuple{Vararg{ExecutionReceipt}})
else
receipt.state = _EXECUTION_RECEIPT_PROVIDER_FAILURE
receipt.failure = provider_failure
_release_receipt_lease!(receipt)
end
end
first_failure === nothing && receipt.failure !== nothing &&
Expand Down
8 changes: 4 additions & 4 deletions src/execution/program_inspection.jl
Original file line number Diff line number Diff line change
Expand Up @@ -690,9 +690,9 @@ function inspect(prepared::PreparedPlan; level = nothing)
drained = prepared.drained,
outstanding = prepared.outstanding,
poisoned = prepared.poisoned,
provider_completions = _lane_wait_count(prepared.lane),
provider_completions = _lane_completion_count(prepared.lane),
provider_scope_completions =
_lane_scope_wait_count(prepared.lane),
_lane_scope_completion_count(prepared.lane),
validation_transfers = _lane_transfer_count(prepared.lane),
),
)
Expand Down Expand Up @@ -772,8 +772,8 @@ function execution_contract(prepared::PreparedPlan)
receipt_scope = _lane_wait_scope(lane),
receipt_cumulative = _lane_cumulative(lane),
receipt_selective = _lane_selective(lane),
observed_provider_completions = _lane_wait_count(lane),
observed_scope_completions = _lane_scope_wait_count(lane),
observed_provider_completions = _lane_completion_count(lane),
observed_scope_completions = _lane_scope_completion_count(lane),
observed_validation_transfers = _lane_transfer_count(lane),
)
end
4 changes: 3 additions & 1 deletion src/execution/stage_program.jl
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,7 @@ end
struct _PreparedStageProgram{E}
launches::Vector{_AbstractPreparedStageLaunch}
execution_gate::E
validation_host::Matrix{UInt32}
end


Expand Down Expand Up @@ -962,7 +963,8 @@ Base.@nospecializeinfer Base.@noinline function _prepare_stage_program(
annotated === error ? rethrow() : throw(annotated)
end
end
return _PreparedStageProgram(launches, workspace.execution_gate)
return _PreparedStageProgram(
launches, workspace.execution_gate, program_host)
end

# `_BoundLaw` already proves that the scientific bindings are mutually legal.
Expand Down
47 changes: 22 additions & 25 deletions src/execution/stage_program_kernelabstractions.jl
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Hardware-agnostic KernelAbstractions provider. Kernel launches are implicitly
# ordered by the backend; one KernelAbstractions.synchronize call is the only
# execution-visibility boundary. No native stream, queue, or event is exposed.
# ordered by the backend. Provider completion is observed either by an explicit
# synchronize or by a nonempty blocking device-to-host validation copy. No
# native stream, queue, or event is exposed.

struct _KernelAbstractionsDeviceGetter end
@inline (::_KernelAbstractionsDeviceGetter)(backend) =
Expand All @@ -14,15 +15,15 @@ mutable struct _KernelAbstractionsScope{B, D, T, E, G}
const device_getter::G
poisoned::Bool
poison_reason::Any
synchronizations::Int
completions::Int
transfers::Int
submitted_ordinal::UInt64
settled_ordinal::UInt64
end

mutable struct _KernelAbstractionsLane{S} <: _AbstractProviderLane
const scope::S
waits::Int
completions::Int
transfers::Int
end

Expand Down Expand Up @@ -117,13 +118,13 @@ _lane_transfer_law(::_KernelAbstractionsLane) = :same_owner_task_only
_lane_cumulative(::_KernelAbstractionsLane) = true
_lane_selective(::_KernelAbstractionsLane) = false
_lane_error_observation(::_KernelAbstractionsLane) = (
synchronization = :kernelabstractions_backend_contract,
completion = :blocking_validation_transfer_or_backend_synchronize,
asynchronous_failures = :backend_defined,
failure_scope = :backend_owner_task,
)
_lane_wait_count(lane::_KernelAbstractionsLane) = lane.waits
_lane_scope_wait_count(lane::_KernelAbstractionsLane) =
lane.scope.synchronizations
_lane_completion_count(lane::_KernelAbstractionsLane) = lane.completions
_lane_scope_completion_count(lane::_KernelAbstractionsLane) =
lane.scope.completions
_lane_transfer_count(lane::_KernelAbstractionsLane) = lane.transfers
_lane_same_wait_scope(
first::_KernelAbstractionsLane, second::_KernelAbstractionsLane
Expand Down Expand Up @@ -191,12 +192,12 @@ _validate_provider_capacity(::_KernelAbstractionsLane, evidence, capacity) =
nothing

function _synchronize_lane_tail!(lane::_KernelAbstractionsLane)
lane.waits += 1
lane.completions += 1
scope = lane.scope
current_task() === scope.owner || throw(LocalMathValidationError(
"KernelAbstractions tail drain must use the preparing owner task"
))
scope.synchronizations += 1
scope.completions += 1
try
KernelAbstractions.synchronize(scope.backend)
catch error
Expand All @@ -209,20 +210,22 @@ function _synchronize_lane_tail!(lane::_KernelAbstractionsLane)
return nothing
end

function _settle_lane_tail!(lane::_KernelAbstractionsLane, statuses::Tuple)
function _settle_lane_tail!(lane::_KernelAbstractionsLane, device, host)
_validate_lane_current!(lane)
lane.waits += 1
lane.completions += 1
scope = lane.scope
scope.synchronizations += 1
scope.completions += 1
try
if isempty(statuses)
if isempty(host)
KernelAbstractions.synchronize(scope.backend)
else
# A host-visible device-to-host copy is the provider completion
# operation. GPU providers must complete their queued prefix before
# returning host data; issuing a second explicit synchronize would
# duplicate that boundary (Metal and CUDA copies are blocking).
_transfer_validation_statuses!(statuses)
# duplicate that boundary. If a provider ever aliases the two
# representations, the copy is not a distinct completion operation.
device === host && KernelAbstractions.synchronize(scope.backend)
_transfer_validation_status!(device, host)
scope.transfers += 1
lane.transfers += 1
end
Expand All @@ -233,11 +236,10 @@ function _settle_lane_tail!(lane::_KernelAbstractionsLane, statuses::Tuple)
return nothing
end

function _transfer_settled_validation_statuses!(
lane::_KernelAbstractionsLane, statuses::Tuple)
isempty(statuses) && return nothing
function _transfer_settled_validation_status!(
lane::_KernelAbstractionsLane, device, host)
try
_transfer_validation_statuses!(statuses)
_transfer_validation_status!(device, host)
catch error
_poison_lane!(lane, error)
rethrow()
Expand All @@ -250,11 +252,6 @@ end
_drain_lane_tail!(lane::_KernelAbstractionsLane) =
_synchronize_lane_tail!(lane)

function _wait_lane!(lane::_KernelAbstractionsLane)
_validate_lane_current!(lane)
return _synchronize_lane_tail!(lane)
end

function _atomic_capability(
backend::KernelAbstractions.Backend,
type::Type,
Expand Down
51 changes: 28 additions & 23 deletions src/execution/validation_support.jl
Original file line number Diff line number Diff line change
Expand Up @@ -94,34 +94,39 @@ _is_publication_validation_error(error) = error isa LocalMathValidationError &&
:runtime_ordered_fold_validation,
)

_prepared_validation_statuses(prepared) = ()
@inline function _prepared_validation_statuses(prepared::PreparedPlan)
# Every contextual Stage status views the same program-level device/host
# buffer. Settlement and success gating therefore need one representative,
# not a flattened tuple whose type grows with total program length.
return (first(prepared.runtime.launches).status,)
end

@inline _prepared_validation_status_groups(prepared::PreparedPlan) =
(map(launch -> launch.status, prepared.runtime.launches),)

function _transfer_validation_statuses!(statuses::Tuple)
isempty(statuses) && return nothing
# Every Stage status is a contextual view of this same program-level
# buffer. One host-visible copy is therefore the complete settlement.
status = first(statuses)
copyto!(status.host, status.device)
@inline function _transfer_validation_status!(device, host)
copyto!(host, device)
return nothing
end

function _prepared_validation_error_at(prepared, lease_index::Int32)
for statuses in _prepared_validation_status_groups(prepared)
for status in statuses
error = _validated_publication_error(status, Int(lease_index))
error === nothing || return error
end
runtime = prepared.runtime
host = runtime.validation_host
failure_class = @inbounds host[_VALIDATION_FAILURE_CLASS, lease_index]
failure_class == UInt32(0) &&
return nothing
recorded_stage = _validation_decode_int32(
@inbounds host[_VALIDATION_STAGE_INDEX, lease_index]
)
# Stage zero is publication suppression propagated from a failed dependency.
# Exact receipt traversal remains the diagnostic and ordering authority.
recorded_stage == 0 && return nothing
for launch in runtime.launches
status = launch.status
status.stage == recorded_stage || continue
return _validated_publication_error(status, Int(lease_index))
end
return nothing
return LocalMathValidationError(
"runtime validation status references an unprepared stage";
stage = :wait,
contract = :validation_status_stage,
expected = :prepared_stage,
actual = (
recorded_stage,
failure_class = _validation_decode_int32(failure_class),
),
hint = "discard the prepared plan and report the invalid validation status",
)
end

function _exact_host_int(value, purpose; stage::Symbol = :plan)
Expand Down
Loading