File size: 48,925 Bytes
6bc44d6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
"""The Spaces' page, four tabs over :mod:`space.pipeline` and the configured source (:mod:`space.sources`).

The acquisition tab: a preset (the source's), custom shells (on a stored pulse timing; at their own δ / Δ / TE, or
from an uploaded Camino scheme, in DiSCo's full mode); the source's tissue-and-scanner panel; the noise; the
tracker; the scanner (ideal, or a catalogued machine whose every catalogued term the replay applies at the phantom's
place in the bore, and whose gradient limit refuses a shell it cannot play; docs/scanner.md); B, the same run with one
knob changed; N
tracker keys for the tractogram's own spread; the run button, whose progress names the stage it is in. The truth
tab: the source's input and ground truth. The Replay DWI Explorer: the ingredient maps of the tiers, the noise-free
layers (the tier ladder, A, B) and their differences in the DWI or in MD / FA, per shell, and the estimated against
the true responses where the source has them. The results tab: a DWI slice with the FODs' principal directions,
the tractograms and connectomes of A and B against the source's truth, the spread over keys, the accuracy and round
trip tables, the timings, the downloads.

The page is built from the source's class (:meth:`space.pipeline.Source.describe`, ``presets``, ``panel``,
``knobs``, ``estimated_seconds``) before any data loads; a run is :func:`prepare_runs` on the host (a lookup of the
source's share that needs no device, in the page's cache), :func:`compute` on the device (which computes the share
the page had not cached first, so the GPU call is entered as soon as the request arrives), :func:`present` back on
the host. Nothing scientific lives here: every number comes from the pipeline or the source, every figure from
:mod:`space.viewers`."""
from __future__ import annotations

import os
import shutil
import tempfile
import threading
import time

import numpy as np
from dataclasses import replace

from . import pipeline as P
from . import sources
from . import viewers as V

MAX_SHELLS = 4
CUSTOM = "custom shells"
UPLOADED = "uploaded scheme"
PER_SHELL = 7                                            # on, timing class, b, directions, delta, Delta, TE
NO_KNOB = P.NO_KNOB
IDEAL = P.IDEAL
RESPONSES_ROW = "worker · the packs' responses not cached on the page"     # the timings row tools/live.py prints
SAMPLE = 10_000                                          # streamlines kept in the page's state and in the sample .tck
OUTPUTS = ("result", "headline", "dwi_view", "tract_view", "mats", "timings", "tck", "volumes", "z_slider", "m_slider",
           "tract_view_b", "mats_b", "b_row", "floor_view", "accuracy", "spread_view",
           "explore_view", "metric_view", "ingredient_pool", "ingredient_contact", "ingredient_field", "layer_table", "layer_choice",
           "ez_slider", "em_slider", "truth_view", "fractions_view", "lobar_view", "roundtrip", "response_view")   # the run button's outputs, in order
EXPLORE_MODES = ("signal", "minus the previous layer", "B − A", "(B − A) ÷ floor")
METRICS = ("DWI", "MD (µm²/ms)", "FA")
_state = {"source": None, "error": None}
_lock = threading.Lock()
_RUNS = tempfile.mkdtemp(prefix="disco-runs-")          # one directory per run tag, replaced on every run


def _load():
    """The source, loaded once per process (a failed load is retried on the next call)."""
    with _lock:
        if _state["source"] is None:
            try:
                t0 = time.perf_counter()
                cfg = P.config()
                src = sources.source(cfg)
                _state["warm"] = src.warm_in_background()           # the brain's warm-up runs beside the serving page
                _state["regions"] = V.region_markers(src.regions)
                _state["gt_views"] = None
                _state["source"] = src; _state["cfg"] = cfg; _state["load_seconds"] = time.perf_counter() - t0; _state["error"] = None
            except Exception as e:                       # shown on the page instead of a dead Space; the next call tries again
                _state["error"] = repr(e)
        return _state


def _protocol_from_inputs(S, cfg, preset, n_b0, *shell_inputs, scheme=None, full=False):
    """The acquisition the page asks for: a preset of the source class ``S``, the shell rows (in full mode each with
    its own delta / Delta / TE in ms, else on a stored timing class), or in full mode an uploaded Camino scheme."""
    if preset == UPLOADED:
        if not full:
            raise ValueError("a scheme upload needs full mode (DISCO_MODE=full with the columnar pack)")
        if not scheme:
            raise ValueError("upload a Camino .scheme file")
        return P.protocol_from_scheme(scheme)
    if preset in S.presets(cfg):
        return S.protocol(cfg, preset)
    if preset != CUSTOM:
        raise ValueError(f"unknown acquisition {preset!r}")
    shells = []
    for k in range(MAX_SHELLS):
        on, shape, b, n, delta, Delta, TE = shell_inputs[PER_SHELL * k: PER_SHELL * (k + 1)]
        if on and full:
            shells.append(P.Shell.free(float(b), int(n), float(delta) * 1e-3, float(Delta) * 1e-3, float(TE) * 1e-3))
        elif on:
            shells.append(P.Shell(str(shape), float(b), int(n)))
    return P.Protocol(tuple(shells), n_b0=int(n_b0), name=CUSTOM)


def physics_values(S, cfg, *values):
    """The panel's inputs as a dict by the source class ``S``'s :attr:`~space.pipeline.Panel.fields`."""
    fields = S.panel(cfg).fields
    if len(values) != len(fields):
        raise ValueError(f"the physics panel has {len(fields)} inputs, got {len(values)}")
    return dict(zip(fields, values))


def gradient_text(cfg, protocol, shapes, scanner):
    """The per-shell gradient table as Markdown, and whether every shell is playable on the menu's ``scanner``."""
    rows = P.playable(protocol, shapes, P.machine(cfg, scanner))
    lines = ["| shell | b (s/mm²) | timing | needs | limit |", "|---|---|---|---|---|"]
    for name, b, delta, Delta, G, G_max, ok in rows:
        limit = "any" if G_max is None else f"{G_max * 1e3:.0f} mT/m {'✓' if ok else '✗ cannot play'}"
        lines.append(f"| {name} | {b:g} | δ {delta * 1e3:g} / Δ {Delta * 1e3:g} ms | {G * 1e3:.0f} mT/m | {limit} |")
    return "\n".join(lines), all(r[-1] for r in rows)


def plan_runs(cfg, source, preset, n_b0, snr_on, snr, scheme_file, knob, scanner, values, shell_inputs):
    """The runs the button asks for, validated before any work: ``[(tag, protocol, snr_on, snr, physics)]`` for A
    and, with a knob, B; refused with the reason when the scanner cannot play a shell or the source cannot do the
    run."""
    S = type(source)
    full = source.mode == "full"
    protocol = _protocol_from_inputs(S, cfg, preset, n_b0, *shell_inputs, scheme=scheme_file, full=full)
    key = P.machine(cfg, scanner)
    runs = [("A", protocol, snr_on, snr, S.physics_from(cfg, values, scanner=key))]
    knobs = S.knobs(cfg)
    if knob not in knobs:
        raise ValueError(f"unknown knob {knob!r}")
    change = knobs[knob]
    if change is not None:
        if key is not None and change[0] in ("field", "b0"):
            raise ValueError(f"the {scanner} fixes the field and its direction: the knob {knob!r} would make B the same run; "
                             "choose the ideal scanner to vary them")
        pb, on_b, snr_b, vb = S.apply_knob(cfg, change, protocol, snr_on, snr, values)
        runs.append(("B", pb, on_b, snr_b, S.physics_from(cfg, vb, scanner=key)))
    for tag, prot, _, _, physics in runs:
        table, ok = gradient_text(cfg, prot, source.shapes, scanner)
        if not ok:
            raise ValueError(f"{scanner} cannot play run {tag}'s shells (square pulses):\n\n{table}")
        source.validate(prot, physics)
    return runs


def _split(S, cfg, rest):
    """The run signature's tail: the physics panel's values (a dict) and the shell rows."""
    n = len(S.panel(cfg).fields)
    return physics_values(S, cfg, *rest[:n]), rest[n:]


def response_entries(source, runs, ladder_on):
    """The :meth:`~space.pipeline.Source.prepare` entries of :func:`plan_runs`'s ``runs``: ``[(slot, meas, physics)]``
    for A, B and, with the ladder, A's rungs; ``slot`` is ``("runs", tag)`` or ``("ladder", i)``."""
    out = [(("runs", tag), P.measurements(prot, source.shapes), ph) for tag, prot, _, _, ph in runs]
    if ladder_on:
        meas_a = out[0][1]
        out += [(("ladder", i), meas_a, ph) for i, (_, ph) in enumerate(source.ladder_steps(runs[0][4]))]
    return out


def prepare_runs(preset, n_b0, snr_on, snr, density, max_angle, step_mm, key, scheme_file, knob, scanner, n_keys, ladder_on, *rest):
    """The source's share of the runs the inputs ask for as the page's process has it cached
    (:meth:`~space.pipeline.Source.cached`), looked up and never computed, so the GPU call follows the request at
    once: ``{"runs": {tag: ...}, "ladder": [...] or None}``, None for an entry not cached (or a source with nothing
    to prepare); :func:`compute` computes the missing ones."""
    state = _load()
    if state["error"]:
        raise ValueError(f"the source did not load: {state['error']}")
    source, cfg = state["source"], state["cfg"]
    values, shell_inputs = _split(type(source), cfg, rest)
    runs = plan_runs(cfg, source, preset, n_b0, snr_on, snr, scheme_file, knob, scanner, values, shell_inputs)
    out = dict(runs={}, ladder=[] if ladder_on else None)
    for (where, k), meas, ph in response_entries(source, runs, ladder_on):
        hit = source.cached(meas, ph)
        if where == "runs":
            out["runs"][k] = hit
        else:
            out["ladder"].append(hit)
    return out


def _fill(source, runs, ladder_on, prepared):
    """:func:`prepare_runs`'s entries the page had not cached, computed here (inside the GPU call), a generator: a
    stage text when there are any, then returns ``(prepared, {response_key: result}, seconds)``: the entries
    complete, the ones computed here by key, the time (None for a source with nothing to prepare)."""
    entries = [(slot, meas, ph) for slot, meas, ph in response_entries(source, runs, ladder_on) if source.response_key(meas, ph) is not None]
    if not entries:
        return prepared, {}, None
    out = dict(runs=dict(prepared["runs"]), ladder=None if prepared["ladder"] is None else list(prepared["ladder"]))
    missing = [(slot, meas, ph) for slot, meas, ph in entries if out[slot[0]][slot[1]] is None]
    if missing:
        yield (f"**computing the packs' responses on the worker: not cached on the page** ({len(missing)} of {len(entries)}) …", 0.0)
    t0 = time.perf_counter()
    computed = {}
    for (where, k), meas, ph in missing:
        r = source.cached(meas, ph)                     # an earlier entry of this run with the same key (B = A but its M0)
        if r is None:
            r = source.prepare(meas, ph)
            computed[source.response_key(meas, ph)] = r
        out[where][k] = r
    return out, computed, time.perf_counter() - t0


def _result_state(res, sample, load_seconds):
    """What the results tab needs, kept per session: float32 volumes, the peaks, the streamline sample, the numbers."""
    pk, amp = P.peaks(res.sh)
    fr = res.extras.get("fractions")
    return dict(dwi=res.dwi.astype(np.float32), floor=res.floor.astype(np.float32), floor_median=P.floor_stats(res)["median"], meas=res.meas,
                peaks=pk.astype(np.float32), peak_amp=amp.astype(np.float32), matrix=res.matrix, score=res.score, seconds=res.seconds,
                load_seconds=load_seconds, tractogram=sample, n_streamlines=len(res.tractogram), shape=res.dwi.shape[:3], name=res.protocol.name,
                extras=dict(fractions=None if fr is None else fr.astype(np.float16)))


def _layer_labels(physics, knob):
    """The explorer's names for A and B: A is the ladder's top rung, named by every tier it has on; B is A with the
    knob."""
    tiers = [q for q in ("relaxation", "contact", "field") if physics and getattr(physics, q)]
    a = "A = bare" + "".join(f" + {q}" for q in tiers) if tiers else "A = bare diffusion"
    return a, f"B = A with {knob}"


def _explorer_state(results, ladder, ingredients, knob, source):
    """What the Replay DWI Explorer keeps per session: the noise-free layers (the ladder, then A, then B) as float16
    volumes with their tensor maps, the replay floor, the ingredient images' layers, the per-shell layer differences."""
    ra = results["A"]; mask = np.isfinite(ra.clean[..., 0])
    name_a, name_b = _layer_labels(ra.physics, knob)
    layers = list(ladder) + [(name_a, ra.clean)] + ([(name_b, results["B"].clean)] if "B" in results else [])
    metrics = {}
    for label, vol in layers:
        try:
            md, fa = P.dti(vol, ra.meas, mask)
        except ValueError:                                  # too few low-b rows for a tensor: no metric maps
            md = fa = None
        metrics[label] = (None if md is None else md.astype(np.float16), None if fa is None else fa.astype(np.float16))
    return dict(layers=[(label, vol.astype(np.float16)) for label, vol in layers], metrics=metrics, floor=ra.floor.astype(np.float32),
                meas=ra.meas, mask=mask, ingredients=source.ingredient_layers(ingredients), differences=P.layer_differences(layers, ra.meas, mask),
                snr=ra.snr, physics=ra.physics)


def _status(text):
    import gradio as gr
    keep = gr.update()
    return (keep, text) + (keep,) * (len(OUTPUTS) - 2)


def _one_run(tag, source, protocol, snr_on, snr, tracking, physics, prepared, t0):
    """One pipeline run as a generator of ``(text, fraction)`` stage updates, then the :class:`Result` last: ``tag``
    is A or B in the stage text."""
    texts = source.describe(source.cfg)["stages"]
    for item in P.run_stages(source, protocol, snr=(float(snr) if snr_on else None), tracking=tracking, physics=physics, prepared=prepared):
        if isinstance(item, P.Result):
            yield item
            return
        stage, k, n = item
        yield (f"**{tag} · {k + 1}/{n} {texts[stage]}** … ({protocol.n_meas} measurements, {time.perf_counter() - t0:.0f} s so far)", (k + 0.5) / (n + 1))


def _write_files(tag, res, source):
    """The run's files under its tag (the previous run's under the same tag replaced): the full tractogram and a
    sample as .tck, the DWI and FOD volumes, the source's own files; ``(sample, {"tck": [...], "volumes": [...]})``."""
    out = os.path.join(_RUNS, tag); shutil.rmtree(out, ignore_errors=True); os.makedirs(out)
    stem = f"{source.describe(source.cfg)['files']}_{tag}_{res.protocol.name.replace(' ', '_')}"
    tck = os.path.join(out, f"{stem}.tck"); res.tractogram.to_tck(tck)
    sample = P.sample_tractogram(res.tractogram, SAMPLE)
    sample_path = os.path.join(out, f"{stem}_sample{SAMPLE // 1000}k.tck"); sample.to_tck(sample_path)
    vols = P.write_volumes(res, out, prefix=stem, affine=source.affine)
    return sample, dict(tck=[tck, sample_path], volumes=[vols["dwi"], vols["bvals"], vols["bvecs"], vols["fod"]] + source.extra_files(res, out, stem))


def explore(ex, layer, mode, metric, z, m):
    """The explorer's map: ``layer`` of the state's layers, in ``mode`` (the signal, its difference to the previous
    layer, B minus A, or that over the replay floor), as the DWI at measurement ``m`` or a tensor metric, slice ``z``."""
    if not ex:
        return None
    labels = [l for l, _ in ex["layers"]]; vols = dict(ex["layers"])
    if layer not in vols:
        layer = labels[-1]
    z = int(z); m = min(int(m), len(ex["meas"].bvals) - 1)
    k = METRICS.index(metric) if metric in METRICS else 0

    def field(label):
        if k == 0:
            return np.asarray(vols[label], np.float32)[..., m]
        md, fa = ex["metrics"][label]
        v = md if k == 1 else fa
        return None if v is None else np.asarray(v, np.float32)
    what = f"{metric} of {layer}" if k else f"{layer}: b = {ex['meas'].bvals[m]:g} s/mm², measurement {m}"
    x = field(layer)
    if x is None:
        return None
    if mode == "signal":
        return V.map_slice(x, z, f"{what}, slice z = {z}", cmap="gray" if k == 0 else "viridis", vmin=0 if k == 0 else None, vmax=1 if k != 1 else None)
    if mode == "minus the previous layer":
        i = labels.index(layer)
        if i == 0:
            return V.map_slice(x, z, f"{layer} is the first layer: its signal, slice z = {z}", cmap="gray", vmin=0, vmax=1)
        prev = field(labels[i - 1])
        return V.map_slice(x - prev, z, f"{what} minus {labels[i - 1]}, slice z = {z}", symmetric=True)
    a_label = next((l for l in labels if l.startswith("A")), None); b_label = next((l for l in labels if l.startswith("B")), None)
    if a_label is None or b_label is None:
        return V.map_slice(x, z, f"no B in this run (choose a knob in the acquisition tab): {what}, slice z = {z}", cmap="gray" if k == 0 else "viridis")
    d = field(b_label) - field(a_label)
    if mode == "B − A":
        return V.map_slice(d, z, f"B − A, {metric} at measurement {m}" if k == 0 else f"B − A, {metric}", symmetric=True)
    if not np.any(ex["floor"] > 0):
        return V.map_slice(d, z, "this source has no per-voxel replay floor: B − A, " + (f"{metric} at measurement {m}" if k == 0 else metric), symmetric=True)
    return V.map_slice(d / np.where(ex["floor"] > 0, ex["floor"], np.nan), z, f"(B − A) ÷ replay floor, {metric}" + (f" at measurement {m}" if k == 0 else ""), symmetric=True)


def ingredient_views(ex, z):
    """The explorer's three ingredient images at slice ``z``, from the source's layers of the run's ingredients."""
    layers = (ex or {}).get("ingredients") or [None, None, None]
    z = int(z)
    return tuple(None if item is None or item[1] is None else V.map_slice(item[1], z, f"{item[0]}, z = {z}", **item[2]) for item in layers)


def layer_table(ex):
    rows = [[a, b, f"{sh:g}", f"{med:.4f}", f"{p99:.4f}"] for a, b, sh, med, p99 in ex["differences"]] if ex else []
    if ex and np.any(ex["floor"] > 0):
        f = ex["floor"][ex["mask"]]
        rows.append(["replay floor", "", "", f"{np.median(f):.4f}", f"{np.quantile(f, 0.99):.4f}"])
    return rows


def compute(prepared, preset, n_b0, snr_on, snr, density, max_angle, step_mm, key, scheme_file, knob, scanner, n_keys, ladder_on, *rest):
    """Everything a run needs the device for, a generator: ``(text, fraction)`` stage updates while it works, then
    the payload last, a dict of plain data (:class:`P.Result` per tag, the explorer's ladder and ingredient maps, the
    tracker-key spread, the seconds of each step, the source's share the page had not cached as computed here by key
    and its seconds, the clock at the handoff). ``prepared`` is :func:`prepare_runs`'s, ``rest`` the physics panel
    then the shell rows. On a shared GPU pool this runs in the forked worker and its yields cross to the page's
    process; nothing here draws or writes a file, so the device is held for the compute alone."""
    state = _load()
    if state["error"]:
        raise ValueError(f"the source did not load: {state['error']}")
    source, cfg = state["source"], state["cfg"]
    S = type(source)
    values, shell_inputs = _split(S, cfg, rest)
    runs = plan_runs(cfg, source, preset, n_b0, snr_on, snr, scheme_file, knob, scanner, values, shell_inputs)
    for tag, prot, _, _, physics in runs:                # full mode: what each run reads, before any byte moves
        plan = source.plan(P.measurements(prot, source.shapes), physics)
        if plan:
            yield (f"**{tag}: full replay of {plan['rows']:,} rows, {plan['bytes'] / 1e9:.1f} GB to read, about "
                   f"{plan['estimated_seconds'] / 60:.0f} min** (bands {plan['K']}, field modes {plan['modes']})", 0.0)
    t0 = time.perf_counter()
    prepared, kernels, response_seconds = yield from _fill(source, runs, ladder_on, prepared)
    tracking = S.tracking(cfg, density=density, max_angle=max_angle, step=step_mm, key=key)
    results = {}; seconds = {}
    try:
        for tag, prot, on, s_, ph in runs:
            for item in _one_run(tag, source, prot, on, s_, tracking, ph, prepared["runs"][tag], t0):
                if isinstance(item, P.Result):
                    results[tag] = item
                else:
                    yield item
        meas_a = results["A"].meas; physics_a = runs[0][4]
        ladder = []
        if ladder_on and physics_a is not None and not physics_a.bare:
            yield (f"**Replay DWI Explorer · the tier ladder of A, noise-free** … ({time.perf_counter() - t0:.0f} s so far)", 0.8)
            t = time.perf_counter(); ladder = source.ladder(meas_a, physics_a, prepared["ladder"]); seconds["explorer · ladder"] = time.perf_counter() - t
        yield (f"**Replay DWI Explorer · the ingredient maps** … ({time.perf_counter() - t0:.0f} s so far)", 0.85)
        t = time.perf_counter(); ingredients = source.ingredients(meas_a, physics_a, prepared["runs"]["A"]); seconds["explorer · ingredients"] = time.perf_counter() - t
        spread = None
        if int(n_keys) > 1:                                 # A's tracking repeated over further keys: the tractogram's own spread
            keys = [int(key) + 1 + i for i in range(int(n_keys) - 1)]
            mats = [results["A"].matrix]; scores = [results["A"].score]
            for i, (k, M, sc, secs) in enumerate(P.repeat_tracking(results["A"], source, tracking, keys)):
                mats.append(M); scores.append(sc)
                yield (f"**A · tracking again with key {k} ({i + 2}/{int(n_keys)})** … ({time.perf_counter() - t0:.0f} s so far)", 0.9)
            spread = P.pair_spread(mats, scores, key=source.score_key())
    finally:
        source.release()
    yield dict(results={tag: _light(r) for tag, r in results.items()}, ladder=[(label, vol.astype(np.float32)) for label, vol in ladder],
               ingredients=ingredients, spread=spread, knob=knob, seconds=seconds, kernels=kernels, response_seconds=response_seconds,
               compute_seconds=time.perf_counter() - t0, handed_off_at=time.time())


def _light(res):
    """``res`` with its DWI volumes in float32 for the handoff: what the page keeps, draws, fits and writes is float32
    or narrower, so the payload carries half the bytes and the compute's own arithmetic is untouched."""
    return replace(res, dwi=res.dwi.astype(np.float32), clean=None if res.clean is None else res.clean.astype(np.float32))


def present(payload, state):
    """The page's outputs (:data:`OUTPUTS`, in order) from a compute payload: the files, the per-session states, the
    headline and the figures (A's source share from the page's cache, which holds it after :func:`run_pipeline`).
    Runs where the page runs, never on the device."""
    import gradio as gr
    received = time.time()
    source = state["source"]; load_seconds = state["load_seconds"]; regions = state["regions"]
    results = payload["results"]; ladder = payload["ladder"]; ingredients = payload["ingredients"]; spread = payload["spread"]; knob = payload["knob"]
    post = {}
    if payload.get("response_seconds") is not None:
        post[RESPONSES_ROW] = payload["response_seconds"]
    post.update(payload["seconds"])
    t = time.perf_counter()
    samples = {}; files = {}
    for tag, r in results.items():
        samples[tag], files[tag] = _write_files(tag, r, source)
    post["page · files"] = time.perf_counter() - t; t = time.perf_counter()
    rs = {tag: _result_state(r, samples[tag], load_seconds) for tag, r in results.items()}
    ex = _explorer_state(results, ladder, ingredients, knob, source)
    post["page · states"] = time.perf_counter() - t; t = time.perf_counter()
    labels = [l for l, _ in ex["layers"]]; name_a = labels[-2] if "B" in results else labels[-1]
    ra = rs["A"]; rb = rs.get("B")
    headline = source.score_text("A", results["A"])
    if rb:
        headline += "<br>" + source.score_text("B", results["B"]) + "<br>" + source.compare_text(source.compare(results["A"], results["B"]), knob)
    if spread:
        headline += (f"<br>**A over {spread['n']} tracker keys: {spread['key']} {spread['pearson_mean']:.3f} ± {spread['pearson_std']:.3f}**, "
                     f"median pair count CV {spread['cv_median']:.2f}, {spread['pairs_always']} pairs in every run, {spread['pairs_any']} in any.")
    d = source.describe(source.cfg)
    headline += " The Replay DWI Explorer is the third tab, the results the fourth."
    z0 = ra["dwi"].shape[2] // 2
    m0 = int(np.flatnonzero(~ra["meas"].b0)[0]) if (~ra["meas"].b0).any() else 0
    timings = [[f"{tag} · {r}", t_] for tag in rs for r, t_ in V.timings_rows(rs[tag]["seconds"], rs[tag]["load_seconds"])] if rb else V.timings_rows(ra["seconds"], ra["load_seconds"])
    timings += [["device held (compute)", f"{payload['compute_seconds']:.2f}"], ["handoff to the page", f"{received - payload['handed_off_at']:.2f}"]]
    timings += [[k, f"{v:.2f}"] for k, v in post.items()]
    rs["explorer"] = ex                                  # the page state: A, B and the explorer's layers
    box = V.grid_box(ra["shape"], source.affine) if d["views"]["truth"] else ra["shape"]
    pa = source.cached(results["A"].meas, results["A"].physics)
    out = (rs, headline, V.dwi_slice(ra["dwi"], ra["meas"], z0, m0, peaks=ra["peaks"], peak_amp=ra["peak_amp"], overlay=True, label="A: "),
           V.tractogram3d(ra["tractogram"], regions, box, total=ra["n_streamlines"]),
           source.matrices(results["A"].matrix, results["A"].score, results["A"].reference),
           timings, [t_ for f in files.values() for t_ in f["tck"]], [v for f in files.values() for v in f["volumes"]],
           _slider_update(z0, ra["dwi"].shape[2] - 1), _slider_update(m0, results["A"].protocol.n_meas - 1),
           V.tractogram3d(rb["tractogram"], regions, box, total=rb["n_streamlines"]) if rb else None,
           source.matrices(results["B"].matrix, results["B"].score, results["B"].reference) if rb else None,
           gr.update(visible=rb is not None),
           V.floor_slice(ra["floor"], z0, ra["floor_median"], label="A: ") if d["views"]["floor"] else None, source.accuracy(results["A"]),
           V.spread_matrices(spread) if spread else None,
           explore(ex, name_a, "minus the previous layer" if len(labels) > 1 else "signal", METRICS[0], z0, m0),
           explore(ex, name_a, "signal", METRICS[2], z0, m0), *ingredient_views(ex, z0), layer_table(ex),
           gr.update(choices=labels, value=name_a), _slider_update(z0, ra["dwi"].shape[2] - 1), _slider_update(m0, results["A"].protocol.n_meas - 1),
           source.truth_view(ra["dwi"], ra["meas"], z0, m0), source.fractions_view(ra["extras"], z0),
           source.lobar(results["A"].matrix, results["A"].score, results["A"].reference), source.roundtrip_rows(results["A"]),
           source.response_view(pa, results["A"].extras))
    timings.append(["page · figures", f"{time.perf_counter() - t:.2f}"])
    return out


def run_pipeline(compute_fn, *args, progress=None):
    """The run button, a generator: the source's share of the runs the page has cached (:func:`prepare_runs`, a
    lookup), then the stage texts of ``compute_fn`` (:func:`compute`, or it wrapped for a GPU pool) into the headline
    as they arrive (the other outputs untouched, the progress bar following), then the share the call computed kept
    in the page's cache and the page's outputs from its payload (:func:`present`)."""
    import gradio as gr
    state = _load()
    if state["error"]:
        raise gr.Error(f"the source did not load: {state['error']}")
    if progress:
        progress(0.0, desc="starting")
    payload = None
    try:
        yield _status("**starting the run** (the packs' responses looked up on the page) …")
        prepared = prepare_runs(*args)
        for item in compute_fn(prepared, *args):
            if isinstance(item, dict):
                payload = item
            else:
                text, fraction = item
                if progress:
                    progress(fraction, desc=text.split("**")[1] if "**" in text else text)
                yield _status(text)
    except (ValueError, KeyError) as e:
        raise gr.Error(str(e))
    if payload.get("kernels"):
        state["source"].keep(payload["kernels"])
    yield _status(f"**drawing** … (device held {payload['compute_seconds']:.0f} s)")
    yield present(payload, state)


GPU_TIERS = (("logged out", 120), ("free account", 300), ("PRO", 2400))    # ZeroGPU's daily quota per visitor tier, seconds


def estimated_seconds(preset, n_b0, snr_on, snr, density, max_angle, step_mm, key, scheme_file, knob, scanner, n_keys, ladder_on, *rest):
    """The GPU seconds a run reserves on the shared pool, from its inputs (the same positional inputs as
    :func:`run_pipeline`): the configured source's measured cost model
    (:meth:`~space.pipeline.Source.estimated_seconds`) with the source's share the call computes because the page has
    it not cached (:func:`uncached_responses`). The pool refuses a request above the visitor's daily quota
    (:data:`GPU_TIERS`) and kills a run that outlives its reservation, so this is the measured cost with its margin,
    not a generous one; 480 when the inputs make no protocol."""
    try:
        cfg = P.config()
        S = sources.source_class(cfg)
        n = len(S.panel(cfg).fields)
        protocol = _protocol_from_inputs(S, cfg, preset, n_b0, *rest[n:], scheme=scheme_file, full=S.mode == "full")
    except Exception:
        return 480
    responses = uncached_responses(preset, n_b0, snr_on, snr, density, max_angle, step_mm, key, scheme_file, knob, scanner, n_keys, ladder_on, *rest)
    try:
        machine = P.machine(cfg, scanner)
    except ValueError:
        return 480
    return S.estimated_seconds(cfg, protocol, density=density, knob=knob, n_keys=n_keys, ladder=ladder_on, responses=responses, scanner=machine)


def uncached_responses(preset, n_b0, snr_on, snr, density, max_angle, step_mm, key, scheme_file, knob, scanner, n_keys, ladder_on, *rest):
    """``(state, n_meas, saves)`` per entry of the source's share of the run the inputs ask for that the page's
    process has not cached (:meth:`~space.pipeline.Source.responses`): what the GPU call computes before its replay. Empty when
    the source is not loaded in this process (the pool's entry loads it before the page serves) or refuses the run
    (it is refused before the GPU call)."""
    source = _state["source"]
    if source is None:
        return []
    cfg = _state["cfg"]
    try:
        values, shell_inputs = _split(type(source), cfg, rest)
        runs = plan_runs(cfg, source, preset, n_b0, snr_on, snr, scheme_file, knob, scanner, values, shell_inputs)
    except (ValueError, KeyError):
        return []
    return source.responses([(meas, ph) for _, meas, ph in response_entries(source, runs, ladder_on)])


def gpu_seconds_text(*args):
    """The readout under the run button on the pool: the seconds this run reserves against the tiers' quotas, and
    which tiers can run it (a request above a visitor's daily quota is refused by the pool before it starts)."""
    secs = estimated_seconds(*args)
    fits = [name for name, cap in GPU_TIERS if secs <= cap]
    who = ("a visitor " + ", ".join(fits)) if fits else "no tier: split the run (drop B, the ladder or the extra keys)"
    caps = ", ".join(f"{name} {cap // 60} min" for name, cap in GPU_TIERS)
    return (f"**This run reserves {secs} s of GPU.** ZeroGPU grants each visitor a daily quota ({caps}) and refuses a "
            f"single request above it (\"larger than the maximum allowed\"): this run can be started by {who}. "
            f"Log in to Hugging Face in this browser to use your own quota.")


def _slider_update(value, maximum):
    import gradio as gr
    return gr.update(value=int(value), maximum=int(maximum))


def redraw_explorer(rs, layer, mode, metric, z, m):
    ex = rs and rs.get("explorer")
    return (explore(ex, layer, mode, metric, z, m), explore(ex, layer, mode, METRICS[2] if metric == METRICS[0] else metric, z, m), *ingredient_views(ex, z))


def redraw_slice(rs, z, m, overlay, which):
    if not rs or which not in rs:
        return None, None, None, None
    r = rs[which]; source = _load()["source"]
    views = source.describe(source.cfg)["views"]
    return (V.dwi_slice(r["dwi"], r["meas"], int(z), int(m), peaks=r["peaks"], peak_amp=r["peak_amp"], overlay=bool(overlay), label=f"{which}: "),
            V.floor_slice(r["floor"], int(z), r["floor_median"], label=f"{which}: ") if views["floor"] else None,
            source.truth_view(r["dwi"], r["meas"], int(z), int(m)), source.fractions_view(r["extras"], int(z)))


def ground_truth_views():
    """The truth tab's views, drawn once per process."""
    state = _load()
    if state["error"]:
        return None, None
    if state["gt_views"] is None:
        state["gt_views"] = state["source"].truth_views()
    return state["gt_views"]


def _widget(c):
    """The Gradio component of a panel :class:`~space.pipeline.Control` (an input kind)."""
    import gradio as gr
    common = dict(label=c.label, visible=c.visible, interactive=c.interactive)
    if c.kind == "checkbox":
        return gr.Checkbox(value=bool(c.value), **common)
    if c.kind == "number":
        return gr.Number(value=c.value, **common)
    if c.kind == "slider":
        return gr.Slider(c.minimum, c.maximum, value=c.value, step=c.step, **common)
    if c.kind == "dropdown":
        return gr.Dropdown(list(c.choices), value=c.value, **common)
    raise ValueError(f"unknown control kind {c.kind!r}")


def build(runner=None, cfg=None):
    """The Blocks of the configured source (``cfg``, else :func:`space.pipeline.config`). ``runner`` wraps
    :func:`compute` for the run button (the ZeroGPU entry passes ``spaces.GPU(...)``): the device part of a run; the
    page looks up its cache before it and draws from its payload after it, in this process."""
    import gradio as gr
    cfg = cfg or P.config()
    S = sources.source_class(cfg)
    d = S.describe(cfg)
    full = S.mode == "full"
    shapes = S.shapes_of(cfg); presets = S.presets(cfg) + [CUSTOM] + ([UPLOADED] if full else [])
    shape_names = [n for n in shapes]
    panel = S.panel(cfg); tc = S.tracking_controls(cfg)
    with gr.Blocks(title=d["title"], delete_cache=(3600, 3600)) as demo:
        gr.Markdown(
            d["heading"]
            + ("\n\n**Before you press run:** the GPU time of a run is charged to *your* Hugging Face quota, not the Space's: "
               "2 minutes a day logged out, 5 with a free account, 40 with PRO. The line under the run button says how many "
               "seconds the configured run reserves and which of those can start it; a request above your quota is refused "
               "before it starts, so log in to Hugging Face in this browser if you want more than one run a day." if runner is not None else "")
            + "\n\n**Fixed here:** " + "; ".join(d["fixed"]) + ".")
        result = gr.State(None)
        with gr.Tabs():
            with gr.Tab("1 · acquisition, tissue and scanner"):
                with gr.Row():
                    with gr.Column(scale=1):
                        preset = gr.Dropdown(presets, value=presets[0], label="acquisition")
                        n_b0 = gr.Slider(1, 10, value=1, step=1, label="b = 0 measurements (custom shells)")
                        shell_inputs = []
                        for k in range(MAX_SHELLS):
                            with gr.Row():
                                on = gr.Checkbox(value=(k == 0), label=f"shell {k + 1}")
                                shape = gr.Dropdown(shape_names, value=shape_names[0], label="pulse timing δ / Δ", visible=not full)
                                b = gr.Number(value=[1000, 2000, 3000, 6000][k], label="b (s/mm²)")
                                n = gr.Slider(6, 128, value=[30, 60, 90, 60][k], step=1, label="directions")
                                delta = gr.Number(value=10.2, label="δ (ms)", visible=full)
                                Delta = gr.Number(value=16.7, label="Δ (ms)", visible=full)
                                TE = gr.Number(value=53.5, label="TE (ms)", visible=full)
                            shell_inputs += [on, shape, b, n, delta, Delta, TE]
                        gr.Markdown(d["acquisition"])
                        scheme_file = gr.File(label="uploaded scheme: Camino STEJSKALTANNER (.scheme)", file_count="single", type="filepath", visible=full)
                    with gr.Column(scale=1):
                        field_presets = {f"{f:g} T": float(f) for f in panel.field_presets}
                        widgets = {}; catalogue_note = reset = field_preset = None
                        with gr.Accordion("tissue and scanner: the physics tiers", open=True):
                            for row in panel.rows:
                                with gr.Row():
                                    for c in row:
                                        if c.kind == "field_preset":
                                            field_preset = gr.Dropdown(list(c.choices), value=c.value, label=c.label)
                                        elif c.kind == "catalogue_note":
                                            catalogue_note = gr.Markdown(c.value)
                                        elif c.kind == "reset":
                                            reset = gr.Button(c.label, size="sm")
                                        elif c.kind == "markdown":
                                            gr.Markdown(c.value)
                                        else:
                                            widgets[c.name] = _widget(c)
                            gr.Markdown(d["tissue"])
                        physics_inputs = [widgets[name] for name in panel.fields]
                        with gr.Row():
                            snr_on = gr.Checkbox(value=True, label="add Rician noise")
                            snr = gr.Slider(5, 100, value=30, step=1, label="SNR at M0 (a full water voxel before relaxation; each voxel's b = 0 SNR follows its tissue)")
                        density = gr.Slider(tc["density"][0], tc["density"][1], value=tc["density"][2], step=tc["density"][3], label=tc["density"][4])
                        max_angle = gr.Slider(tc["max_angle"][0], tc["max_angle"][1], value=tc["max_angle"][2], step=tc["max_angle"][3], label=tc["max_angle"][4])
                        step_mm = gr.Slider(tc["step"][0], tc["step"][1], value=tc["step"][2], step=tc["step"][3], label=tc["step"][4])
                        key = gr.Number(value=0, precision=0, label="random key")
                        knob = gr.Dropdown(list(S.knobs(cfg)), value=NO_KNOB, label="B: the same run with one knob changed")
                        scanner = gr.Dropdown([IDEAL] + list(P.machines(cfg)), value=IDEAL,
                                              label="scanner: a catalogued machine sets the field and its direction, plays every term "
                                                    "its catalogue entry carries, and refuses a shell beyond its gradient limit")
                        scanner_terms = gr.Markdown(P.scanner_text(cfg, IDEAL, S.refused_machines(cfg)))
                        gradients = gr.Markdown()
                        n_keys = gr.Slider(1, 8, value=1, step=1, label="repeat A's tracking over N keys (the tractogram's own spread)")
                        ladder_on = gr.Checkbox(value=not full, visible=not full, label=d["ladder"])
                        go = gr.Button(d["run_label"], variant="primary")
                        gpu_text = gr.Markdown(visible=runner is not None)
                        headline = gr.Markdown()
            with gr.Tab(d["truth_tab"]) as gt_tab:
                strands_view = gr.Plot(label=d["truth_labels"][0])
                gt_matrix = gr.Image(label=d["truth_labels"][1], type="pil")
                gr.Markdown(f"**The truth:** {d['truth']}.")
            with gr.Tab("3 · Replay DWI Explorer"):
                gr.Markdown(d["explorer"])
                with gr.Row():
                    ingredient_pool = gr.Image(label=d["ingredients"][0], type="pil")
                    ingredient_contact = gr.Image(label=d["ingredients"][1], type="pil")
                    ingredient_field = gr.Image(label=d["ingredients"][2], type="pil")
                with gr.Row():
                    layer_choice = gr.Dropdown(["A"], value="A", label="layer")
                    mode = gr.Radio(list(EXPLORE_MODES), value=EXPLORE_MODES[0], label="show")
                    metric = gr.Radio(list(METRICS), value=METRICS[0], label="quantity")
                with gr.Row():
                    ez_slider = gr.Slider(0, 39, value=20, step=1, label="axial slice z")
                    em_slider = gr.Slider(0, 363, value=0, step=1, label="measurement (DWI only)")
                with gr.Row():
                    explore_view = gr.Image(label="the chosen layer and mode", type="pil")
                    metric_view = gr.Image(label="the same for FA (or the chosen metric)", type="pil")
                layer_tbl = gr.Dataframe(headers=["from", "to", "b (s/mm²)", "median |ΔS|", "99 % |ΔS|"], label="layer differences per shell, and the replay floor", interactive=False)
                response_view = gr.Image(label="estimated vs true response per tissue", type="pil", visible=d["views"]["response"])
            with gr.Tab(d["results_tab"]):
                with gr.Row():
                    with gr.Column(scale=1):
                        dwi_view = gr.Image(label="DWI slice", type="pil")
                        with gr.Row():
                            z_slider = gr.Slider(0, 39, value=20, step=1, label="axial slice z")
                            m_slider = gr.Slider(0, 363, value=0, step=1, label="measurement")
                            overlay = gr.Checkbox(value=True, label="FOD principal directions")
                            which = gr.Radio(["A", "B"], value="A", label="run")
                    with gr.Column(scale=1, visible=d["views"]["truth"]):
                        truth_view = gr.Image(label="the same slice with the input FOD's principal directions", type="pil")
                    with gr.Column(scale=1):
                        tract_view = gr.Plot(label="tractogram A")
                fractions_view = gr.Image(label="recovered vs input fractions", type="pil", visible=d["views"]["fractions"])
                mats = gr.Image(label="connectome A vs its truth", type="pil")
                lobar_view = gr.Image(label="connectome A vs its truth over the lobar groups", type="pil", visible=d["views"]["lobar"])
                with gr.Row(visible=False) as b_row:
                    tract_view_b = gr.Plot(label="tractogram B")
                    mats_b = gr.Image(label="connectome B vs its truth", type="pil")
                spread_view = gr.Image(label="A over N tracker keys: mean and spread per pair", type="pil")
                roundtrip = gr.Dataframe(headers=["what", "value"], label="the round trip: the reconstruction and the connectome against the input",
                                         interactive=False, wrap=True, visible=d["views"]["roundtrip"])
                with gr.Accordion("accuracy: what this replay is an approximation of", open=False):
                    gr.Markdown(d["accuracy"])
                    with gr.Row():
                        floor_view = gr.Image(label="split-half floor (this run, the slice above)", type="pil", visible=d["views"]["floor"])
                        accuracy = gr.Dataframe(headers=["what", "value"], label="the source and this run", interactive=False, wrap=True)
                with gr.Row():
                    timings = gr.Dataframe(headers=["stage", "seconds"], label="timings", interactive=False)
                    with gr.Column():
                        tck = gr.File(label="tractograms (.tck, MRtrix): every streamline, and a 10k sample", file_count="multiple")
                        volumes = gr.File(label=d["volumes_label"], file_count="multiple")
        compute_fn = compute if runner is None else runner(compute)      # the device part alone runs under a pool's GPU

        def run_with_progress(*args, progress=gr.Progress()):
            yield from run_pipeline(compute_fn, *args, progress=progress)
        tissue_numbers = [widgets[name] for name in panel.catalogue] + [catalogue_note]
        field_preset.change(lambda name: [field_presets[name]] + S.catalogue_numbers(cfg, field_presets[name]), inputs=field_preset,
                            outputs=[widgets["field_T"]] + tissue_numbers, show_progress="hidden")
        reset.click(lambda f: S.catalogue_numbers(cfg, f), inputs=widgets["field_T"], outputs=tissue_numbers, show_progress="hidden")

        def show_gradients(preset_, n_b0_, scanner_, scheme_, *shells_):
            try:
                prot = _protocol_from_inputs(S, cfg, preset_, n_b0_, *shells_, scheme=scheme_, full=full)
            except (ValueError, KeyError) as e:
                return f"({e})"
            return gradient_text(cfg, prot, shapes, scanner_)[0]

        def choose_scanner(label, field_now):
            """The menu's machine: its term table, its field on the slider and the catalogue's tissue at it (the ideal
            scanner keeps the field as set)."""
            key = P.machine(cfg, label)
            f = float(field_now) if key is None else float(P.limits(key).field_T)
            return [P.scanner_text(cfg, label, S.refused_machines(cfg)), f] + (S.catalogue_numbers(cfg, f) if key is not None else [gr.update()] * len(tissue_numbers))
        scanner.change(choose_scanner, inputs=[scanner, widgets["field_T"]], outputs=[scanner_terms, widgets["field_T"]] + tissue_numbers,
                       show_progress="hidden")
        for ctl in (preset, scanner, scheme_file, *shell_inputs):
            ctl.change(show_gradients, inputs=[preset, n_b0, scanner, scheme_file, *shell_inputs], outputs=gradients, show_progress="hidden")
        outputs = dict(result=result, headline=headline, dwi_view=dwi_view, tract_view=tract_view, mats=mats, timings=timings, tck=tck,
                       volumes=volumes, z_slider=z_slider, m_slider=m_slider, tract_view_b=tract_view_b, mats_b=mats_b, b_row=b_row,
                       floor_view=floor_view, accuracy=accuracy, spread_view=spread_view, explore_view=explore_view, metric_view=metric_view,
                       ingredient_pool=ingredient_pool, ingredient_contact=ingredient_contact, ingredient_field=ingredient_field,
                       layer_table=layer_tbl, layer_choice=layer_choice, ez_slider=ez_slider, em_slider=em_slider, truth_view=truth_view,
                       fractions_view=fractions_view, lobar_view=lobar_view, roundtrip=roundtrip, response_view=response_view)
        for ctl in (layer_choice, mode, metric, ez_slider, em_slider):
            ctl.change(redraw_explorer, inputs=[result, layer_choice, mode, metric, ez_slider, em_slider],
                       outputs=[explore_view, metric_view, ingredient_pool, ingredient_contact, ingredient_field], show_progress="hidden")
        run_inputs = [preset, n_b0, snr_on, snr, density, max_angle, step_mm, key, scheme_file, knob, scanner, n_keys, ladder_on, *physics_inputs, *shell_inputs]
        if runner is not None:                                   # the pool: what the run reserves, live as the inputs change
            gr.on([c.change for c in run_inputs] + [demo.load], gpu_seconds_text, inputs=run_inputs, outputs=gpu_text, show_progress="hidden")
        run_event = go.click(run_with_progress, inputs=run_inputs,
                             outputs=[outputs[name] for name in OUTPUTS], concurrency_limit=1, api_name="run_pipeline")   # the endpoint tools/live.py drives
        if runner is not None:                                   # the run's responses are now cached on the page: the reservation falls
            run_event.then(gpu_seconds_text, inputs=run_inputs, outputs=gpu_text, show_progress="hidden")
        for ctl in (z_slider, m_slider, overlay, which):
            ctl.change(redraw_slice, inputs=[result, z_slider, m_slider, overlay, which], outputs=[dwi_view, floor_view, truth_view, fractions_view], show_progress="hidden")
        gt_tab.select(ground_truth_views, outputs=[strands_view, gt_matrix])
        demo.load(lambda: (_load().get("error") and f"**the source did not load:** {_load()['error']}") or "", outputs=headline)
    return demo


if __name__ == "__main__":
    threading.Thread(target=_load, daemon=True).start()          # download and warm while the page comes up
    build().queue(max_size=8).launch(server_name="0.0.0.0", server_port=int(os.environ.get("PORT", 7860)))