Skip to content

Commit 8c71c62

Browse files
meggartlazarusA
andauthored
Move Dagger to a DiskArrayEninge extension (#52)
* merge changes * Move Dagger to extension * update docs deps --------- Co-authored-by: Lazaro Alonso <lazarus.alon@gmail.com>
1 parent df01be1 commit 8c71c62

5 files changed

Lines changed: 163 additions & 144 deletions

File tree

Project.toml

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,6 @@ version = "0.2.6"
44
authors = ["Fabian Gans <fgans@bgc-jena.mpg.de> and contributors"]
55

66
[deps]
7-
Dagger = "d58978e5-989f-55fb-8d15-ea34adc7bf54"
87
DiskArrays = "3c3547ce-8d99-4f5e-a174-61eb10b00ae3"
98
Distributed = "8ba89e20-285c-5b6f-9357-94700520ee1b"
109
FileWatching = "7b1f6079-737a-58dc-b8bc-7a2ca5c1b5ee"
@@ -23,9 +22,15 @@ Statistics = "10745b16-79ce-11e8-11f9-7d13ad32a3b2"
2322
StatsBase = "2913bbd2-ae8a-5f71-8c99-4fb6c76f3a91"
2423
Zarr = "0a941bbe-ad1d-11e8-39d9-ab76183a1d99"
2524

25+
[weakdeps]
26+
Dagger = "d58978e5-989f-55fb-8d15-ea34adc7bf54"
27+
28+
[extensions]
29+
DaggerExt = "Dagger"
30+
2631
[compat]
2732
Dagger = "0.18, 0.19"
28-
DiskArrays = "0.3, 0.4.10"
33+
DiskArrays = "0.4.10"
2934
Graphs = "1"
3035
Interpolations = "0.14, 0.15, 0.16"
3136
Ipopt = "1"

docs/Project.toml

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,4 +4,8 @@ Documenter = "e30172f5-a6a5-5a46-863b-614d45cd2de4"
44
DocumenterVitepress = "4710194d-e776-4893-9690-8d956a29c365"
55

66
[sources]
7-
DiskArrayEngine = {path = ".."}
7+
DiskArrayEngine = {path = ".."}
8+
9+
[compat]
10+
Documenter = "1"
11+
DocumenterVitepress = "0.3"

docs/package.json

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,8 @@
1010
"markdown-it-footnote": "^4.0.0",
1111
"markdown-it-mathjax3": "^4.3.2",
1212
"vitepress": "^1.6.3",
13-
"vitepress-plugin-tabs": "^0.6.0"
13+
"@mdit/plugin-mathjax": "^0.25.0",
14+
"@mdit/plugin-tex": "^0.23.1",
15+
"vitepress-plugin-tabs": "^0.6.0"
1416
}
1517
}

ext/DaggerExt.jl

Lines changed: 147 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,147 @@
1+
module DaggerExt
2+
import Dagger
3+
import DiskArrayEngine: DaggerRunner, create_outars, plan_to_loopranges, generate_outbuffers, generate_inbuffers, merge_outbuffer_collection,
4+
run_loop, read_range, extract_outbuffer, run_block, put_buffer, clean_aggregator, merge_all_outbuffers, flush_all_outbuffers,
5+
get_procgroups, DiskEngineScheduler, run_group, subset_loopranges, callback_and_runfilters
6+
import Distributed: myid, Distributed
7+
8+
make_outbuffer_shard(op,runnerloopranges,workerthreads) = Dagger.shard(per_thread=workerthreads) do
9+
generate_outbuffers(op.outspecs,op.f,runnerloopranges)
10+
end
11+
function DaggerRunner(op,exec_plan,outars=create_outars(op,exec_plan;par_only=true);
12+
workerthreads=false,threaded=true,showprogress=true,restartfile=nothing,restartmode=:continue)
13+
inars = op.inars
14+
loopranges = plan_to_loopranges(exec_plan)
15+
inbuffers = Dagger.shard(per_thread=workerthreads) do
16+
generate_inbuffers(inars, loopranges)
17+
end
18+
try
19+
generate_outbuffers(op.outspecs,op.f, loopranges)
20+
catch e
21+
rethrow(e)
22+
end
23+
cb, runfilter = callback_and_runfilters(loopranges, showprogress, restartfile, restartmode)
24+
cbc = if !isempty(cb)
25+
channel = Distributed.RemoteChannel(() -> Channel{Union{Nothing,eltype(loopranges)}}(), 1)
26+
@async while true
27+
update = take!(channel)
28+
isnothing(update) && break
29+
notify_callback(cb, update)
30+
end
31+
channel
32+
else
33+
()
34+
end
35+
DaggerRunner(op,loopranges,outars, threaded, workerthreads, inbuffers, cbc, runfilter)
36+
end
37+
38+
buffer_mergefunc(red,::Type{<:Union{Dagger.Chunk,Dagger.Thunk, Dagger.EagerThunk}}) = (buf1,buf2) -> begin
39+
@debug "Creating Dagger merge function"
40+
@debug "Fetching x"
41+
fx = fetch(buf1)
42+
@debug "Fetching y"
43+
fy = fetch(buf2)
44+
@debug "Calling function"
45+
merge_outbuffer_collection.(fx,fy,(red,))
46+
end
47+
48+
function run_loop(runner::DaggerRunner,loopranges,outbuffers...;groupspecs=nothing)
49+
@noinline run_loop(runner,runner.op,runner.inbuffers_pure,runner.loopranges,runner.workerthreads,
50+
runner.outars,runner.threaded,loopranges,runner.cbc,runner.runfilter,outbuffers...;groupspecs)
51+
end
52+
53+
54+
function run_loop(::DaggerRunner,op,inbuffers_pure,runnerloopranges,workerthreads,
55+
outars,threaded, loopranges,cbc,runfilter,outbuffers...;groupspecs=nothing)
56+
@debug "Groupspecs are ", groupspecs
57+
piddir = if groupspecs !== nothing && any(i->in(:output_chunk,i.reasons),groupspecs)
58+
tempname()
59+
else
60+
nothing
61+
end
62+
@debug "Pidddir is $piddir"
63+
local_outbuffers = make_outbuffer_shard(op,runnerloopranges,workerthreads)
64+
op = op
65+
r = broadcast(loopranges) do inow
66+
Dagger.spawn(inbuffers_pure,local_outbuffers,inow,piddir,outars,cbc,runfilter) do inbuffers_pure, outbuffers, inow, piddir, outars,cbc,runfilter
67+
#default_loopbody(inow, op, inbuffers_pure, outbuffers, threaded, outars, cbc, runfilter, piddir)
68+
@debug myid(), " Starting block ", inow
69+
inbuffers_wrapped = read_range.((inow,),op.inars,inbuffers_pure);
70+
outbuffers_now = extract_outbuffer.((inow,),op.outspecs,op.f.init,op.f.buftype,outbuffers)
71+
run_block(op,inow,inbuffers_wrapped,outbuffers_now,threaded)
72+
@debug myid(), "Finished running block ", inow
73+
put_buffer.((inow,),outbuffers_now, outars, (piddir,))
74+
clean_aggregator.(outbuffers)
75+
true
76+
end
77+
end
78+
all(fetch.(r)) || error("Some workers errored")
79+
@debug myid(), " Fetched everything"
80+
if (groupspecs !== nothing) && any(i->in(:reducedim,i.reasons),groupspecs)
81+
@debug "Merging buffers"
82+
procs = unique(Dagger.processor.(fetch.(r,raw=true)))
83+
@debug "Affected processors are $procs"
84+
buffers_used = collect(v for (k,v) in local_outbuffers.chunks if any(p->matches_proc(k,p),procs))
85+
@debug "Merging buffers from $(length(buffers_used)) workers."
86+
buffers_used = fetch.(buffers_used)
87+
collections_merged = merge_all_outbuffers(buffers_used,op.f.red)
88+
@debug "Writing merged buffers $(typeof(collections_merged))"
89+
unflushed_buffers = Dagger.spawn(collections_merged,outars,piddir) do cm,outars,pdir
90+
flush_all_outbuffers(cm,outars,pdir)
91+
end
92+
if !isempty(outbuffers)
93+
outbuffers = last(outbuffers)
94+
@debug "Putting back flushed buffers"
95+
r = Dagger.spawn(unflushed_buffers,outbuffers,red) do rembuf,outbuf, red
96+
foreach(rembuf,outbuf) do r,o
97+
if !isempty(r.buffers)
98+
@debug "Putting back unflushed data"
99+
newagg = merge_outbuffer_collection(o,r,red)
100+
empty!(o)
101+
for k in keys(newagg)
102+
o[k] = newagg
103+
end
104+
end
105+
end
106+
end
107+
fetch(r)
108+
else
109+
@debug "Outbuffers are empty"
110+
return fetch(unflushed_buffers)
111+
@debug "Fetched unflushed buffers"
112+
end
113+
end
114+
GC.gc()
115+
true
116+
end
117+
118+
matches_proc(k::Dagger.ThreadProc,c::Dagger.ThreadProc) = k==c
119+
matches_proc(k::Dagger.OSProc,c::Dagger.ThreadProc) = c.tid != 1 ? error("Processing was not on tid 1") : c.owner == k.pid
120+
121+
function Base.run(runner::DaggerRunner)
122+
@debug "Starting to run"
123+
groups = get_procgroups(runner.op, runner.loopranges, fetch.(runner.outars))
124+
sch = DiskEngineScheduler(groups, runner.loopranges, runner)
125+
opts = runner.workerthreads ? (;) : (;scope = Dagger.scope(thread=1))
126+
Dagger.with_options(;opts...) do
127+
@debug "Calling first run_group"
128+
run_group(sch,nothing)
129+
end
130+
runner.outars
131+
end
132+
133+
function schedule(sch::DiskEngineScheduler,r::DaggerRunner,loopdims,loopsub,groupspecs)
134+
@debug "Starting to schedule: "
135+
r = map(loopsub) do i
136+
lrsub = subset_loopranges(sch.loopranges,loopdims,i.I)
137+
@debug "New split loopranges are: ", lrsub.members
138+
schsub = DiskEngineScheduler(sch.groups,lrsub,sch.runner)
139+
@debug "Spawning"
140+
outbuffers = make_outbuffer_shard(r.op,r.loopranges,r.workerthreads)
141+
Dagger.spawn(schsub,groupspecs,outbuffers) do sched, gs, ob
142+
run_group(sched,gs,ob)
143+
end
144+
end
145+
wait.(r)
146+
end
147+
end

src/daggerrunner.jl

Lines changed: 1 addition & 140 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,3 @@
1-
using Dagger: Dagger, shard
2-
31
struct DaggerRunner
42
op
53
loopranges
@@ -10,142 +8,5 @@ struct DaggerRunner
108
cbc
119
runfilter
1210
end
13-
make_outbuffer_shard(op,runnerloopranges,workerthreads) = Dagger.shard(per_thread=workerthreads) do
14-
generate_outbuffers(op.outspecs,op.f,runnerloopranges)
15-
end
16-
function DaggerRunner(op,exec_plan,outars=create_outars(op,exec_plan;par_only=true);
17-
workerthreads=false,threaded=true,showprogress=true,restartfile=nothing,restartmode=:continue)
18-
inars = op.inars
19-
loopranges = plan_to_loopranges(exec_plan)
20-
inbuffers = Dagger.shard(per_thread=workerthreads) do
21-
generate_inbuffers(inars, loopranges)
22-
end
23-
try
24-
generate_outbuffers(op.outspecs,op.f, loopranges)
25-
catch e
26-
rethrow(e)
27-
end
28-
cb, runfilter = callback_and_runfilters(loopranges, showprogress, restartfile, restartmode)
29-
cbc = if !isempty(cb)
30-
channel = Distributed.RemoteChannel(() -> Channel{Union{Nothing,eltype(loopranges)}}(), 1)
31-
@async while true
32-
update = take!(channel)
33-
isnothing(update) && break
34-
notify_callback(cb, update)
35-
end
36-
channel
37-
else
38-
()
39-
end
40-
DaggerRunner(op,loopranges,outars, threaded, workerthreads, inbuffers, cbc, runfilter)
41-
end
42-
43-
buffer_mergefunc(red,::Type{<:Union{Dagger.Chunk,Dagger.Thunk, Dagger.EagerThunk}}) = (buf1,buf2) -> begin
44-
@debug "Creating Dagger merge function"
45-
@debug "Fetching x"
46-
fx = fetch(buf1)
47-
@debug "Fetching y"
48-
fy = fetch(buf2)
49-
@debug "Calling function"
50-
merge_outbuffer_collection.(fx,fy,(red,))
51-
end
52-
53-
function run_loop(runner::DaggerRunner,loopranges,outbuffers...;groupspecs=nothing)
54-
@noinline run_loop(runner,runner.op,runner.inbuffers_pure,runner.loopranges,runner.workerthreads,
55-
runner.outars,runner.threaded,loopranges,runner.cbc,runner.runfilter,outbuffers...;groupspecs)
56-
end
57-
58-
59-
function run_loop(::DaggerRunner,op,inbuffers_pure,runnerloopranges,workerthreads,
60-
outars,threaded, loopranges,cbc,runfilter,outbuffers...;groupspecs=nothing)
61-
@debug "Groupspecs are ", groupspecs
62-
piddir = if groupspecs !== nothing && any(i->in(:output_chunk,i.reasons),groupspecs)
63-
tempname()
64-
else
65-
nothing
66-
end
67-
@debug "Pidddir is $piddir"
68-
local_outbuffers = make_outbuffer_shard(op,runnerloopranges,workerthreads)
69-
op = op
70-
r = broadcast(loopranges) do inow
71-
Dagger.spawn(inbuffers_pure,local_outbuffers,inow,piddir,outars,cbc,runfilter) do inbuffers_pure, outbuffers, inow, piddir, outars,cbc,runfilter
72-
#default_loopbody(inow, op, inbuffers_pure, outbuffers, threaded, outars, cbc, runfilter, piddir)
73-
@debug myid(), " Starting block ", inow
74-
inbuffers_wrapped = read_range.((inow,),op.inars,inbuffers_pure);
75-
outbuffers_now = extract_outbuffer.((inow,),op.outspecs,op.f.init,op.f.buftype,outbuffers)
76-
run_block(op,inow,inbuffers_wrapped,outbuffers_now,threaded)
77-
@debug myid(), "Finished running block ", inow
78-
put_buffer.((inow,),outbuffers_now, outars, (piddir,))
79-
clean_aggregator.(outbuffers)
80-
true
81-
end
82-
end
83-
all(fetch.(r)) || error("Some workers errored")
84-
@debug myid(), " Fetched everything"
85-
if (groupspecs !== nothing) && any(i->in(:reducedim,i.reasons),groupspecs)
86-
@debug "Merging buffers"
87-
procs = unique(Dagger.processor.(fetch.(r,raw=true)))
88-
@debug "Affected processors are $procs"
89-
buffers_used = collect(v for (k,v) in local_outbuffers.chunks if any(p->matches_proc(k,p),procs))
90-
@debug "Merging buffers from $(length(buffers_used)) workers."
91-
buffers_used = fetch.(buffers_used)
92-
collections_merged = merge_all_outbuffers(buffers_used,op.f.red)
93-
@debug "Writing merged buffers $(typeof(collections_merged))"
94-
unflushed_buffers = Dagger.spawn(collections_merged,outars,piddir) do cm,outars,pdir
95-
flush_all_outbuffers(cm,outars,pdir)
96-
end
97-
if !isempty(outbuffers)
98-
outbuffers = last(outbuffers)
99-
@debug "Putting back flushed buffers"
100-
r = Dagger.spawn(unflushed_buffers,outbuffers,red) do rembuf,outbuf, red
101-
foreach(rembuf,outbuf) do r,o
102-
if !isempty(r.buffers)
103-
@debug "Putting back unflushed data"
104-
newagg = merge_outbuffer_collection(o,r,red)
105-
empty!(o)
106-
for k in keys(newagg)
107-
o[k] = newagg
108-
end
109-
end
110-
end
111-
end
112-
fetch(r)
113-
else
114-
@debug "Outbuffers are empty"
115-
return fetch(unflushed_buffers)
116-
@debug "Fetched unflushed buffers"
117-
end
118-
end
119-
GC.gc()
120-
true
121-
end
122-
123-
matches_proc(k::Dagger.ThreadProc,c::Dagger.ThreadProc) = k==c
124-
matches_proc(k::Dagger.OSProc,c::Dagger.ThreadProc) = c.tid != 1 ? error("Processing was not on tid 1") : c.owner == k.pid
125-
126-
function Base.run(runner::DaggerRunner)
127-
@debug "Starting to run"
128-
groups = get_procgroups(runner.op, runner.loopranges, fetch.(runner.outars))
129-
sch = DiskEngineScheduler(groups, runner.loopranges, runner)
130-
opts = runner.workerthreads ? (;) : (;scope = Dagger.scope(thread=1))
131-
Dagger.with_options(;opts...) do
132-
@debug "Calling first run_group"
133-
run_group(sch,nothing)
134-
end
135-
runner.outars
136-
end
13711

138-
function schedule(sch::DiskEngineScheduler,r::DaggerRunner,loopdims,loopsub,groupspecs)
139-
@debug "Starting to schedule: "
140-
r = map(loopsub) do i
141-
lrsub = subset_loopranges(sch.loopranges,loopdims,i.I)
142-
@debug "New split loopranges are: ", lrsub.members
143-
schsub = DiskEngineScheduler(sch.groups,lrsub,sch.runner)
144-
@debug "Spawning"
145-
outbuffers = make_outbuffer_shard(r.op,r.loopranges,r.workerthreads)
146-
Dagger.spawn(schsub,groupspecs,outbuffers) do sched, gs, ob
147-
run_group(sched,gs,ob)
148-
end
149-
end
150-
wait.(r)
151-
end
12+
DaggerRunner(args...; kwargs...) = error("The Dagger extension for DiskArrayEngine is not loaded. Please activate it by running `import Dagger` in your code.")

0 commit comments

Comments
 (0)