rhbch's picture
arxiv-Link
8d756ce
Raw
History Blame Contribute Delete
3.17 kB
import gradio as gr
from schwarz_demo import DemoConfig, architecture_figure, preview, train, training_figure, operators_figure, solution_figure, summary
def cfg(n, o, lr, cr, local, coarse, glob, assembly, steps):
return DemoConfig(
int(n), int(o), int(lr), int(cr), bool(local), bool(coarse), bool(glob),
assembly == "RAS-style output weighting (experimental)", int(steps)
)
def update(*args):
try:
c = cfg(*args)
return preview(c), architecture_figure(c)
except Exception as e:
return f"⚠️ {e}", None
def run(*args, progress=gr.Progress()):
c = cfg(*args)
r = train(c, lambda s, n: progress(s / n, desc=f"Training {s}/{n}"))
return training_figure(r), operators_figure(r), solution_figure(r), summary(r)
with gr.Blocks(title="Overlapping Schwarz Attention") as demo:
gr.Markdown(
"# 🧩 Overlapping Schwarz Attention Playground: [https://arxiv.org/abs/2606.18525](https://arxiv.org/abs/2606.18525)\n"
"Explore **local-only**, **coarse-only**, and **two-level** attention, "
"with an optional nearest-parameter global low-rank baseline."
)
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("### Schwarz components")
local = gr.Checkbox(True, label="Local overlapping attention blocks")
coarse = gr.Checkbox(True, label="Coarse attention block")
glob = gr.Checkbox(True, label="Train nearest-parameter global low-rank baseline")
n = gr.Slider(2, 16, 8, step=1, label="Subdomains")
o = gr.Slider(0, 8, 2, step=1, label="Overlap")
lr = gr.Slider(1, 24, 4, step=1, label="Local attention rank")
cr = gr.Slider(1, 16, 7, step=1, label="Coarse attention rank")
steps = gr.Dropdown([500, 1000, 2000], value=500, label="Training steps")
with gr.Accordion("Advanced options", open=False):
assembly = gr.Radio(
[
"Symmetric partition-of-unity weighting (paper default)",
"RAS-style output weighting (experimental)",
],
value="Symmetric partition-of-unity weighting (paper default)",
label="Local assembly",
)
gr.Markdown(
"The RAS-style option is an exploratory comparison. "
"The paper uses symmetric partition-of-unity weighting."
)
button = gr.Button("Train and compare", variant="primary")
with gr.Column(scale=2):
info = gr.Markdown()
arch = gr.Plot()
gr.Markdown("### Training and operators")
curve = gr.Plot()
ops = gr.Plot()
sol = gr.Plot()
gr.Markdown("### Results and diagnostics")
table = gr.Markdown()
inputs = [n, o, lr, cr, local, coarse, glob, assembly, steps]
for x in inputs:
x.change(update, inputs, [info, arch])
button.click(run, inputs, [curve, ops, sol, table])
demo.load(update, inputs, [info, arch])
if __name__ == "__main__":
demo.queue(default_concurrency_limit=1).launch()