disco-zero / space /app.py
rfick's picture
Squash history: one commit holding the current files
4b23c8f
Raw History Blame Contribute Delete
48.9 kB
"""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)))