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
239 changes: 204 additions & 35 deletions src/integralop.jl
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,165 @@ for `I` and `J` permutations of `1:numfunctions(test_space)` and
"""
function blockassembler end

"""
AssemblyScratch{Z,T,Q}

Reusable local buffers for one block-assembly operation.

An `AssemblyScratch` must be owned by at most one assembly task at a time.
"""
struct AssemblyScratch{Z,T,Q}
zlocal::Z
tadjq::T
qbuffer::Q
end

"""
AbstractAssemblyScratchSource

Abstract selector for the source of block-assembly scratch buffers.
"""
abstract type AbstractAssemblyScratchSource end

"""
NewAssemblyScratch()

Select fresh buffers for an assembly operation.

`assemblyscratch` allocates a new `AssemblyScratch` when this source is used.
"""
struct NewAssemblyScratch <: AbstractAssemblyScratchSource end

"""
StoredAssemblyScratch(scratch)

Select an existing scratch bundle for an assembly operation.

The referenced bundle is returned without allocation. It must not be used by
another task until the current operation has finished.
"""
struct StoredAssemblyScratch{S} <: AbstractAssemblyScratchSource
scratch::S
end

"""
AssemblyScratchPool(available, available_condition)

Thread-safe pool of reusable scratch bundles for block assembly.

Each bundle is exclusively borrowed by one task. The condition protects the
available-bundle list and lets tasks wait when all bundles are in use.
"""
struct AssemblyScratchPool{S}
available::Vector{S}
available_condition::Threads.Condition
end

"""
acquirescratch(pool::AssemblyScratchPool)

Acquire an available scratch bundle from `pool`.

The call waits if every bundle is currently in use. The caller owns the
returned bundle until it passes it to [`releasescratch!`](@ref).
"""
function acquirescratch(pool::AssemblyScratchPool)
lock(pool.available_condition)
try
while isempty(pool.available)
wait(pool.available_condition)
end
return pop!(pool.available)
finally
unlock(pool.available_condition)
end
end

"""
releasescratch!(pool::AssemblyScratchPool, scratch)

Return `scratch` to `pool` and notify one waiting task.

Every successful call to [`acquirescratch`](@ref) must be paired with one
call to this function, including when assembly exits through an exception.
"""
function releasescratch!(pool::AssemblyScratchPool, scratch)
lock(pool.available_condition)
try
push!(pool.available, scratch)
notify(pool.available_condition)
finally
unlock(pool.available_condition)
end
return nothing
end

"""
assemblyscratch(
::NewAssemblyScratch,
biop,
test_space,
trial_space,
test_assembly_data,
trial_assembly_data,
quadstrat,
)

Allocate an `AssemblyScratch` for the supplied assembly configuration.

The local matrix, trial-adjoint workspace, and quadrature buffer are sized
from the assembly data, spaces, and quadrature strategy.
"""
function assemblyscratch(
::NewAssemblyScratch,
biop,
test_space,
trial_space,
test_assembly_data,
trial_assembly_data,
quadstrat,
)
return AssemblyScratch(
zeros(
scalartype(biop, test_space, trial_space),
size(test_assembly_data.data, 2),
size(trial_assembly_data.data, 2),
),
Vector{eltype(trial_assembly_data.data)}(
undef, size(trial_assembly_data.data, 1)
),
quadraturebuffer(quadstrat, test_space, trial_space),
)
end

"""
assemblyscratch(
source::StoredAssemblyScratch,
biop,
test_space,
trial_space,
test_assembly_data,
trial_assembly_data,
quadstrat,
)

Return the stored scratch bundle without allocating new buffers.

The caller is responsible for ensuring that the bundle is not used
concurrently by another assembly task.
"""
function assemblyscratch(
source::StoredAssemblyScratch,
biop,
test_space,
trial_space,
test_assembly_data,
trial_assembly_data,
quadstrat,
)
return source.scratch
end


# """
# quadrule(operator,test_refspace,trial_refspace,p,test_element,q_trial_element, qd)
Expand Down Expand Up @@ -139,9 +298,6 @@ function assemblechunk!(biop::IntegralOperator, tfs::Space, bfs::Space, store;
tshapes = refspace(tfs); #num_tshapes = numfunctions(tshapes, tdom)
bshapes = refspace(bfs); #num_bshapes = numfunctions(bshapes, bdom)

num_tshapes = size(tad.data, 2)
num_bshapes = size(bad.data, 2)

qs = if CompScienceMeshes.refines(tgeo, bgeo)
TestRefinesTrialQStrat(quadstrat)
elseif CompScienceMeshes.refines(bgeo, tgeo)
Expand All @@ -151,7 +307,6 @@ function assemblechunk!(biop::IntegralOperator, tfs::Space, bfs::Space, store;
end

qd = quaddata(biop, tshapes, bshapes, test_elements, bsis_elements, qs)
zlocal = zeros(scalartype(biop, tfs, bfs), 2num_tshapes, 2num_bshapes)
# @show "after" qs
# assemblechunk_body!(biop,
# tfs, test_elements, tad, tcells,
Expand All @@ -163,7 +318,7 @@ function assemblechunk!(biop::IntegralOperator, tfs::Space, bfs::Space, store;
assemblechunk_body!(biop, tfs, bfs,
test_elements, tcells, tad, eachindex(tcells),
bsis_elements, bcells, bad, eachindex(bcells),
qd, zlocal, store; quadstrat=qs)
qd, store; quadstrat=qs)

end

Expand Down Expand Up @@ -198,19 +353,28 @@ end
function assemblechunk_body!(biop, test_space, trial_space,
test_elements, test_element_ptrs, test_assembly_data, active_test_els,
trial_elements, trial_element_ptrs, trial_assembly_data, active_trial_els,
qd, zlocal, store; quadstrat, scheduler=:serial)
qd, store;
quadstrat,
scheduler=:serial,
scratchsource=NewAssemblyScratch())

num_tshapes = size(test_assembly_data.data,2)
num_bshapes = size(trial_assembly_data.data,2)

@tasks for p in eachindex(active_test_els)
@set scheduler = scheduler
@local begin
zlocal = zeros(scalartype(biop, test_space, trial_space), num_tshapes, num_bshapes)
tadjq = Vector{eltype(trial_assembly_data.data)}(undef, size(trial_assembly_data.data,1))
qbuffer = quadraturebuffer(quadstrat, test_space, trial_space)
end
qaction = ApplyIntegrate(qbuffer)
@local scratch = assemblyscratch(
scratchsource,
biop,
test_space,
trial_space,
test_assembly_data,
trial_assembly_data,
quadstrat,
)
zlocal = scratch.zlocal
tadjq = scratch.tadjq
qaction = ApplyIntegrate(scratch.qbuffer)
P = active_test_els[p]
tcell = test_elements[P]
tptr = test_element_ptrs[P]
Expand Down Expand Up @@ -288,7 +452,7 @@ struct AssembleblockbodyFunctor{B,T1,T2,T3,T4,T5,T6,T7,T8,T9}
trialelements::T5
trialassemblydata::T6
quadraturedata::T7
zlocals::T8
scratchpool::T8
quadstrat::T9
end

Expand All @@ -308,18 +472,25 @@ function (f::AssembleblockbodyFunctor)(testids, trialids, store)
test_element_ptrs = eachindex(f.testelements)
trial_element_ptrs = eachindex(f.trialelements)

assemblechunk_body!(f.biop, f.tfs, f.bfs,
f.testelements, test_element_ptrs, tad1, active_test_els,
f.trialelements, trial_element_ptrs, bad1, active_trial_els,
f.quadraturedata, nothing, store;
quadstrat=f.quadstrat, scheduler=:serial)
scratch = acquirescratch(f.scratchpool)
try
assemblechunk_body!(f.biop, f.tfs, f.bfs,
f.testelements, test_element_ptrs, tad1, active_test_els,
f.trialelements, trial_element_ptrs, bad1, active_trial_els,
f.quadraturedata, store;
quadstrat=f.quadstrat,
scheduler=:serial,
scratchsource=StoredAssemblyScratch(scratch))
finally
releasescratch!(f.scratchpool, scratch)
end
# assembleblock_body!(f.biop, f.tfs, testids, f.testelements, f.testassemblydata,
# f.bfs, trialids, f.trialelements, f.trialassemblydata,
# f.quadraturedata, f.zlocals, store; quadstrat=f.quadstrat)
end

function blockassembler(biop::IntegralOperator, tfs::Space, bfs::Space;
quadstrat=defaultquadstrat(biop, tfs, bfs))
quadstrat=defaultquadstrat(biop, tfs, bfs), primer=nothing)

tgeo = geometry(tfs)
bgeo = geometry(bfs)
Expand All @@ -332,9 +503,13 @@ function blockassembler(biop::IntegralOperator, tfs::Space, bfs::Space;
quadstrat
end

if primer === nothing
primer = assembleblock_primer(biop, tfs, bfs; quadstrat=qs)
end

test_elements, test_assembly_data,
trial_elements, trial_assembly_data,
quadrature_data, zlocals = assembleblock_primer(biop, tfs, bfs; quadstrat=qs)
quadrature_data, scratchpool = primer

return AssembleblockbodyFunctor(
biop,
Expand All @@ -345,7 +520,7 @@ function blockassembler(biop::IntegralOperator, tfs::Space, bfs::Space;
trial_elements,
trial_assembly_data,
quadrature_data,
zlocals,
scratchpool,
qs,
)

Expand Down Expand Up @@ -405,24 +580,18 @@ function assembleblock_primer(biop, tfs, bfs;
test_elements, tad = assemblydata(tfs; onlyactives=false)
bsis_elements, bad = assemblydata(bfs; onlyactives=false)

tgeo = geometry(tfs)
bgeo = geometry(bfs)

tdom = domain(chart(tgeo, first(tgeo)))
bdom = domain(chart(bgeo, first(bgeo)))

tshapes = refspace(tfs); num_tshapes = numfunctions(tshapes, tdom)
bshapes = refspace(bfs); num_bshapes = numfunctions(bshapes, bdom)
tshapes = refspace(tfs)
bshapes = refspace(bfs)

qd = quaddata(biop, tshapes, bshapes, test_elements, bsis_elements, quadstrat)

zlocals = Channel{Matrix{scalartype(biop, tfs, bfs)}}(2*Threads.nthreads())

for _ in 1:2*Threads.nthreads()
put!(zlocals, zeros(scalartype(biop, tfs, bfs), num_tshapes, num_bshapes))
end
scratches = [
assemblyscratch(NewAssemblyScratch(), biop, tfs, bfs, tad, bad, quadstrat)
for _ in 1:(2 * Threads.nthreads())
]
scratchpool = AssemblyScratchPool(scratches, Threads.Condition())

return test_elements, tad, bsis_elements, bad, qd, zlocals
return test_elements, tad, bsis_elements, bad, qd, scratchpool
end

# function assembleblock_body!(biop::IntegralOperator,
Expand Down
2 changes: 1 addition & 1 deletion src/operator.jl
Original file line number Diff line number Diff line change
Expand Up @@ -274,7 +274,7 @@ function assemble!(operator::Operator, testfunctions::Space, trialfunctions::Spa
assemblechunk_body!(operator, testfunctions, trialfunctions,
testelements, eachindex(testelements), testad1, coloredtestelements,
trialelements, eachindex(trialelements), trialad, coloredtrialelements,
qdata, nothing, store; quadstrat=qs, scheduler)
qdata, store; quadstrat=qs, scheduler)
next!(pbar; step = length(testelementcolors[i]) * length(trialelementcolors[j]))
end end
finish!(pbar)
Expand Down
1 change: 1 addition & 0 deletions test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ include("test_local_storage.jl")
include("test_embedding.jl")

include("test_assemblerow.jl")
include("test_blockassembler.jl")
# include("test_mixed_blkassm.jl")
include("test_local_assembly.jl")
include("test_assemble_refinements.jl")
Expand Down
Loading
Loading