mgtotaro commited on
Commit
c74a121
·
1 Parent(s): 31be27c

rendering fix

Browse files
Files changed (2) hide show
  1. app.py +14 -9
  2. data.py +12 -3
app.py CHANGED
@@ -8,6 +8,7 @@ MODELS = ModelFactory.models()
8
 
9
  def app(seq, sub, model_name, acc):
10
  "Main application function"
 
11
  scoring = "masked-marginals" if acc else "wt-marginals"
12
 
13
  # Validate the input
@@ -20,15 +21,17 @@ def app(seq, sub, model_name, acc):
20
  try:
21
  data = Data(seq, sub, model_name, scoring).calculate(progress)
22
  if isinstance(data.image, str):
23
- return ( Image(value=data.image, type='filepath', visible=True)
24
- , HTML()
25
  , DownloadButton(value=data.csv, visible=True) )
26
  else:
27
  return ( Image(visible=False)
28
  , HTML(value=data.image.to_html(), visible=True)
29
  , DownloadButton(value=data.csv, visible=True) )
30
  except Exception as e:
31
- raise Error(str(e))
 
 
32
 
33
  # Create the Gradio interface
34
  with Blocks() as demo:
@@ -38,22 +41,24 @@ with Blocks() as demo:
38
  Markdown(open("header.md", 'r', encoding="utf-8").read())
39
  seq = Textbox( lines=2
40
  , label="Sequence"
41
- , placeholder="FASTA sequence here..."
42
  , value='' )
43
  sub = Textbox( lines=1
44
  , label="Substitutions"
45
- , placeholder="Substitutions here..."
46
  , value='' )
47
  model_name = Dropdown(MODELS, label="Model", value="facebook/esm2_t30_150M_UR50D")
48
  acc_box = Checkbox(value=True, label="Use higher accuracy scoring", interactive=True)
49
  run_btn = Button(value="Run", variant="primary")
50
  dl_btn = DownloadButton(label="Download raw data", visible=False)
 
51
  progress = Progress()
52
- out_html = HTML()
53
  out_img = Image(visible=False)
54
- run_btn.click( fn=app
55
- , inputs=[seq, sub, model_name, acc_box]
56
- , outputs=[out_img, out_html, dl_btn] )
 
 
57
  ex = Examples(
58
  examples=[
59
  [ "MVEQYLLEAIVRDARDGITISDCSRPDNPLVFVNDAFTRMTGYDAEEVIGKNCRFLQRGDINLSAVHTIKIAMLTHEPCLVTLKNYRKDGTIFWNELSLTPIINKNGLITHYLGIQKDVSAQVILNQTLHEENHLLKSNKEMLEYLVNIDALTGLHNRRFLEDQLVIQWKLASRHINTITIFMIDIDYFKAFNDTYGHTAGDEALRTIAKTLNNCFMRGSDFVARYGGEEFTILAIGMTELQAHEYSTKLVQKIENLNIHHKGSPLGHLTISLGYSQANPQYHNDQNLVIEQADRALYSAKVEGKNRAVAYREQ"
 
8
 
9
  def app(seq, sub, model_name, acc):
10
  "Main application function"
11
+ global progress
12
  scoring = "masked-marginals" if acc else "wt-marginals"
13
 
14
  # Validate the input
 
21
  try:
22
  data = Data(seq, sub, model_name, scoring).calculate(progress)
23
  if isinstance(data.image, str):
24
+ return ( Image(value=data.image, type='filepath', height=None, visible=True)
25
+ , HTML(value='')
26
  , DownloadButton(value=data.csv, visible=True) )
27
  else:
28
  return ( Image(visible=False)
29
  , HTML(value=data.image.to_html(), visible=True)
30
  , DownloadButton(value=data.csv, visible=True) )
31
  except Exception as e:
32
+ raise Error(str(e))
33
+ finally:
34
+ progress = Progress()
35
 
36
  # Create the Gradio interface
37
  with Blocks() as demo:
 
41
  Markdown(open("header.md", 'r', encoding="utf-8").read())
42
  seq = Textbox( lines=2
43
  , label="Sequence"
44
+ , placeholder="FASTA sequence here…"
45
  , value='' )
46
  sub = Textbox( lines=1
47
  , label="Substitutions"
48
+ , placeholder="Substitutions here…"
49
  , value='' )
50
  model_name = Dropdown(MODELS, label="Model", value="facebook/esm2_t30_150M_UR50D")
51
  acc_box = Checkbox(value=True, label="Use higher accuracy scoring", interactive=True)
52
  run_btn = Button(value="Run", variant="primary")
53
  dl_btn = DownloadButton(label="Download raw data", visible=False)
54
+ out_html = HTML(visible=False)
55
  progress = Progress()
 
56
  out_img = Image(visible=False)
57
+ run_btn.click( fn=lambda : (HTML(value=None, visible=False), Image(value=None, visible=True, height=128))
58
+ , outputs=[out_html, out_img]
59
+ ).then( fn=app
60
+ , inputs=[seq, sub, model_name, acc_box]
61
+ , outputs=[out_img, out_html, dl_btn] )
62
  ex = Examples(
63
  examples=[
64
  [ "MVEQYLLEAIVRDARDGITISDCSRPDNPLVFVNDAFTRMTGYDAEEVIGKNCRFLQRGDINLSAVHTIKIAMLTHEPCLVTLKNYRKDGTIFWNELSLTPIINKNGLITHYLGIQKDVSAQVILNQTLHEENHLLKSNKEMLEYLVNIDALTGLHNRRFLEDQLVIQWKLASRHINTITIFMIDIDYFKAFNDTYGHTAGDEALRTIAKTLNNCFMRGSDFVARYGGEEFTILAIGMTELQAHEYSTKLVQKIENLNIHHKGSPLGHLTISLGYSQANPQYHNDQNLVIEQADRALYSAKVEGKNRAVAYREQ"
data.py CHANGED
@@ -1,8 +1,9 @@
1
- from math import ceil
2
  import matplotlib.pyplot as plt
3
  import pandas as pd
4
  from re import match
5
  import seaborn as sns
 
6
 
7
  from model import ModelFactory
8
 
@@ -178,12 +179,20 @@ class Data:
178
  ax[i].set_yticklabels(ax[i].get_yticklabels(), rotation=0)
179
  ax[i].set_xticklabels(ax[i].get_xticklabels(), rotation=90)
180
  fig.tight_layout()
181
-
182
  def calculate(self, progress):
183
  "run model and parse output"
 
 
 
 
 
 
 
 
184
  self.progress = progress
185
  self.model.run_model(self)
186
- self.parse_output()
187
  return self
188
 
189
  @property
 
1
+ from math import ceil, exp
2
  import matplotlib.pyplot as plt
3
  import pandas as pd
4
  from re import match
5
  import seaborn as sns
6
+ import asyncio
7
 
8
  from model import ModelFactory
9
 
 
179
  ax[i].set_yticklabels(ax[i].get_yticklabels(), rotation=0)
180
  ax[i].set_xticklabels(ax[i].get_xticklabels(), rotation=90)
181
  fig.tight_layout()
182
+
183
  def calculate(self, progress):
184
  "run model and parse output"
185
+ async def _parse_output():
186
+ self.progress(0, desc="Rendering")
187
+ ren = asyncio.create_task(asyncio.to_thread(self.parse_output))
188
+ t = 0
189
+ while not ren.done():
190
+ await asyncio.sleep(0.33)
191
+ t += 1
192
+ self.progress((1 - exp(-0.1 * t)), desc="Rendering")
193
  self.progress = progress
194
  self.model.run_model(self)
195
+ asyncio.run(_parse_output())
196
  return self
197
 
198
  @property