Spaces:
Running on Zero
Running on Zero
Download space/app.py from rfick/disco-zero: direct link, hf CLI and curl.
- Browser
- Download file 48.9 kB
-
https://huggingface.co/spaces/rfick/disco-zero/resolve/main/space/app.py
- Command line
-
hf download hf://spaces/rfick/disco-zero/space/app.py
-
curl -L -o app.py https://huggingface.co/spaces/rfick/disco-zero/resolve/main/space/app.py
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))) | |