WinslowFan Claude Opus 5.5 commited on
Commit
a8d1387
Β·
1 Parent(s): dd079af

Precompute inner graphs and drop duplicate edge rows

Browse files

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

backend/app/protein/api/routes/explore.py CHANGED
@@ -133,9 +133,11 @@ async def get_inner_graph(
133
  max_states: int = Query(7, ge=1, le=40),
134
  ):
135
  """Inner layer for one sequence: observations of its largest states and their pairwise similarities."""
136
- data = _index_required(
137
- await run_in_threadpool(catalog.inner_graph_nodes, sequence_id.upper(), max_nodes, max_states)
138
- )
 
 
139
  if not data:
140
  raise HTTPException(status_code=404, detail=f"Sequence {sequence_id} not found")
141
  chains = [(n["pdb_id"], n["auth_asym_id"]) for n in data["nodes"]]
 
133
  max_states: int = Query(7, ge=1, le=40),
134
  ):
135
  """Inner layer for one sequence: observations of its largest states and their pairwise similarities."""
136
+ sid = sequence_id.upper()
137
+ pre = await run_in_threadpool(catalog.precomputed_inner_graph, sid, max_nodes, max_states)
138
+ if pre:
139
+ return pre
140
+ data = _index_required(await run_in_threadpool(catalog.inner_graph_nodes, sid, max_nodes, max_states))
141
  if not data:
142
  raise HTTPException(status_code=404, detail=f"Sequence {sequence_id} not found")
143
  chains = [(n["pdb_id"], n["auth_asym_id"]) for n in data["nodes"]]
backend/app/protein/catalog.py CHANGED
@@ -11,6 +11,7 @@ from __future__ import annotations
11
  import json
12
  import re
13
  import sqlite3
 
14
  from functools import lru_cache
15
  from typing import Any, Dict, List, Optional, Tuple
16
 
@@ -408,9 +409,51 @@ def outer_graph(hsets: List[str], per_group: int = 14, include: List[str] = ())
408
  return {"groups": groups, "names": names["homologies"]}
409
 
410
 
411
- def inner_graph_nodes(sequence_id: str, max_nodes: int = 60, max_states: int = 7) -> Optional[Dict[str, Any]]:
412
- """Pick up to max_nodes observations from the sequence's max_states largest states."""
 
 
 
413
  conn = _connect()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
414
  if conn is None:
415
  return None
416
  try:
@@ -435,7 +478,18 @@ def inner_graph_nodes(sequence_id: str, max_nodes: int = 60, max_states: int = 7
435
  )
436
  ]
437
  finally:
438
- conn.close()
 
 
 
 
 
 
 
 
 
 
 
439
 
440
  order = [s["state_id"] for s in states[:max_states]]
441
  pool = [m for m in members if m["state_id"] in set(order)]
@@ -457,11 +511,4 @@ def inner_graph_nodes(sequence_id: str, max_nodes: int = 60, max_states: int = 7
457
  if not added:
458
  break
459
  depth += 1
460
- shown_states = {s["state_id"] for s in states[:MAX_HEATMAP_STATES]}
461
- return {
462
- "sequence": _sequence_dict(row),
463
- "nodes": chosen,
464
- "n_total": len(members),
465
- "states": states,
466
- "state_pairs": [p for p in pairs if p["a"] in shown_states and p["b"] in shown_states],
467
- }
 
11
  import json
12
  import re
13
  import sqlite3
14
+ import zlib
15
  from functools import lru_cache
16
  from typing import Any, Dict, List, Optional, Tuple
17
 
 
409
  return {"groups": groups, "names": names["homologies"]}
410
 
411
 
412
+ FIDELITY_NAMES = ["identical", "low", "medium", "high", None]
413
+
414
+
415
+ def precomputed_inner_graph(sequence_id: str, max_nodes: int, max_states: int) -> Optional[Dict[str, Any]]:
416
+ """Inner graph stored by build_index.py (one row read), or None if not precomputed."""
417
  conn = _connect()
418
+ if conn is None:
419
+ return None
420
+ try:
421
+ try:
422
+ row = conn.execute(
423
+ "SELECT nodes, edges FROM inner_graph WHERE sequence_id = ? AND max_nodes = ? AND max_states = ?",
424
+ (sequence_id, max_nodes, max_states),
425
+ ).fetchone()
426
+ except sqlite3.OperationalError: # older index without the table
427
+ return None
428
+ if row is None:
429
+ return None
430
+ base = inner_graph_nodes(sequence_id, max_nodes, max_states, conn=conn, select=False)
431
+ finally:
432
+ conn.close()
433
+ by_key = {(m["pdb_id"], m["auth_asym_id"]): m for m in base["members"]}
434
+ nodes = [by_key[(p, c)] for p, c in json.loads(row["nodes"]) if (p, c) in by_key]
435
+ edges = [
436
+ {"s": i, "t": j, "sim": None if sim < 0 else sim / 1000, "fid": FIDELITY_NAMES[f]}
437
+ for i, j, sim, f in json.loads(zlib.decompress(row["edges"]))
438
+ ]
439
+ base.pop("members")
440
+ base.update(nodes=nodes, edges=edges)
441
+ return base
442
+
443
+
444
+ def inner_graph_nodes(
445
+ sequence_id: str,
446
+ max_nodes: int = 60,
447
+ max_states: int = 7,
448
+ conn: Optional[sqlite3.Connection] = None,
449
+ select: bool = True,
450
+ ) -> Optional[Dict[str, Any]]:
451
+ """Pick up to max_nodes observations from the sequence's max_states largest states.
452
+
453
+ The selection rule is mirrored by build_index.select_inner_nodes; keep them in sync.
454
+ """
455
+ own = conn is None
456
+ conn = conn or _connect()
457
  if conn is None:
458
  return None
459
  try:
 
478
  )
479
  ]
480
  finally:
481
+ if own:
482
+ conn.close()
483
+
484
+ shown_states = {s["state_id"] for s in states[:MAX_HEATMAP_STATES]}
485
+ base = {
486
+ "sequence": _sequence_dict(row),
487
+ "n_total": len(members),
488
+ "states": states,
489
+ "state_pairs": [p for p in pairs if p["a"] in shown_states and p["b"] in shown_states],
490
+ }
491
+ if not select:
492
+ return {**base, "members": members}
493
 
494
  order = [s["state_id"] for s in states[:max_states]]
495
  pool = [m for m in members if m["state_id"] in set(order)]
 
511
  if not added:
512
  break
513
  depth += 1
514
+ return {**base, "nodes": chosen}
 
 
 
 
 
 
 
backend/app/protein/records.py CHANGED
@@ -186,7 +186,12 @@ def get_transitions(pdb_id: str, chain: str, limit: int = 20000) -> List[Dict[st
186
  conn.close()
187
 
188
  out = []
 
189
  for r in rows:
 
 
 
 
190
  sim = num(r["pair_similarity"])
191
  out.append({
192
  "pdb_id": r["pdb_id_B"].lower(),
@@ -230,6 +235,7 @@ def pair_similarities(chains: List[tuple]) -> List[Dict[str, Any]]:
230
  """Undirected similarity edges among a set of (pdb_id, chain) observations of one sequence."""
231
  keep = {(p.lower(), c): i for i, (p, c) in enumerate(chains)}
232
  out = []
 
233
  conn = _connect()
234
  try:
235
  for i, (pdb, chain) in enumerate(chains):
@@ -239,7 +245,8 @@ def pair_similarities(chains: List[tuple]) -> List[Dict[str, Any]]:
239
  (pdb, chain),
240
  ):
241
  j = keep.get((r[0].lower(), r[1]))
242
- if j is not None and j > i:
 
243
  out.append({"s": i, "t": j, "sim": num(r[2]), "fid": text(r[3])})
244
  finally:
245
  conn.close()
 
186
  conn.close()
187
 
188
  out = []
189
+ seen = set()
190
  for r in rows:
191
+ key = (r["pdb_id_B"].lower(), r["auth_asym_id_B"])
192
+ if key in seen: # the edge table holds some exact duplicate rows
193
+ continue
194
+ seen.add(key)
195
  sim = num(r["pair_similarity"])
196
  out.append({
197
  "pdb_id": r["pdb_id_B"].lower(),
 
235
  """Undirected similarity edges among a set of (pdb_id, chain) observations of one sequence."""
236
  keep = {(p.lower(), c): i for i, (p, c) in enumerate(chains)}
237
  out = []
238
+ seen = set()
239
  conn = _connect()
240
  try:
241
  for i, (pdb, chain) in enumerate(chains):
 
245
  (pdb, chain),
246
  ):
247
  j = keep.get((r[0].lower(), r[1]))
248
+ if j is not None and j > i and (i, j) not in seen:
249
+ seen.add((i, j))
250
  out.append({"s": i, "t": j, "sim": num(r[2]), "fid": text(r[3])})
251
  finally:
252
  conn.close()
backend/scripts/build_index.py CHANGED
@@ -35,6 +35,9 @@ from pathlib import Path
35
  FUNCTION_TEXT_TOP_N = 3
36
  FUNCTION_TEXT_MAX_CHARS = 600
37
  MAX_PAIR_STATES = 40 # state_pair rows are kept only among a sequence's largest states
 
 
 
38
 
39
 
40
  def log(msg: str) -> None:
@@ -512,6 +515,100 @@ def order_stats(src: sqlite3.Connection, limit: int | None) -> dict:
512
  }
513
 
514
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
515
  # ─────────────────────────────────────────────────────────────── overview
516
  def outer_graph_stats(dst: sqlite3.Connection, n_sequences: int) -> dict:
517
  rows = dst.execute(
@@ -607,6 +704,8 @@ def main():
607
  ap.add_argument("--out-dir", type=Path, required=True)
608
  ap.add_argument("--limit", type=int, default=None, help="debug: only read the first N rows")
609
  ap.add_argument("--ecod", type=Path, default=None, help="ecod.latest.domains.txt (for ECOD names)")
 
 
610
  args = ap.parse_args()
611
 
612
  args.out_dir.mkdir(parents=True, exist_ok=True)
@@ -617,6 +716,17 @@ def main():
617
  src = sqlite3.connect(f"file:{args.db}?mode=ro", uri=True)
618
  src.execute("PRAGMA temp_store = FILE")
619
  src.execute("PRAGMA cache_size = -2000000")
 
 
 
 
 
 
 
 
 
 
 
620
  dst = sqlite3.connect(tmp_path)
621
  dst.execute("PRAGMA journal_mode = OFF")
622
  dst.execute("PRAGMA synchronous = OFF")
@@ -637,6 +747,8 @@ def main():
637
  edges = edge_histograms(src, args.limit)
638
  log("pass 4/4: order / disorder labels")
639
  order = order_stats(src, args.limit)
 
 
640
 
641
  overview = build_overview(seqs, ov, n_nodes, edges, state_fid)
642
  overview["outer_graph"] = outer_graph_stats(dst, len(seqs))
 
35
  FUNCTION_TEXT_TOP_N = 3
36
  FUNCTION_TEXT_MAX_CHARS = 600
37
  MAX_PAIR_STATES = 40 # state_pair rows are kept only among a sequence's largest states
38
+ INNER_MAX_NODES = 60 # precomputed inner graph: observations per sequence …
39
+ INNER_MAX_STATES = 7 # … drawn from its largest states (must match app/protein/catalog.py)
40
+ FIDELITY_CODE = {"identical": 0, "low": 1, "medium": 2, "high": 3}
41
 
42
 
43
  def log(msg: str) -> None:
 
515
  }
516
 
517
 
518
+ def select_inner_nodes(members: list, states: list, max_nodes: int, max_states: int) -> list:
519
+ """Round-robin over the largest states. Mirrors catalog.inner_graph_nodes exactly."""
520
+ order = [st for st, _ in states[:max_states]]
521
+ keep = set(order)
522
+ pool = [m for m in members if m[0] in keep]
523
+ if len(pool) <= max_nodes:
524
+ return pool
525
+ by_state: dict = {}
526
+ for m in pool:
527
+ by_state.setdefault(m[0], []).append(m)
528
+ chosen, depth = [], 0
529
+ while len(chosen) < max_nodes:
530
+ added = False
531
+ for st in order:
532
+ bucket = by_state.get(st, [])
533
+ if depth < len(bucket) and len(chosen) < max_nodes:
534
+ chosen.append(bucket[depth])
535
+ added = True
536
+ if not added:
537
+ break
538
+ depth += 1
539
+ return chosen
540
+
541
+
542
+ def build_inner_graphs(src: sqlite3.Connection, dst: sqlite3.Connection) -> None:
543
+ """Precompute each sequence's inner-graph sample and its pairwise similarities.
544
+
545
+ Serving these live needs one edge-index lookup per node, which is slow when the
546
+ DB sits on a network mount; stored here it is a single row read.
547
+ """
548
+ import zlib
549
+
550
+ states: dict = defaultdict(list)
551
+ for seq, st, n in dst.execute(
552
+ "SELECT sequence_id, state_id, n_members FROM state ORDER BY sequence_id, n_members DESC, state_id"
553
+ ):
554
+ states[seq].append((st, n))
555
+ members: dict = defaultdict(list)
556
+ for seq, st, pdb, chain in dst.execute(
557
+ "SELECT sequence_id, state_id, pdb_id, auth_asym_id FROM member "
558
+ "ORDER BY sequence_id, state_id, resolution IS NULL, resolution"
559
+ ):
560
+ members[seq].append((st, pdb.lower(), chain))
561
+
562
+ src.execute("CREATE TEMP TABLE sel (pdb TEXT, chain TEXT, seq TEXT, idx INTEGER, PRIMARY KEY (pdb, chain))")
563
+ nodes_by_seq = {}
564
+ rows = []
565
+ for seq, mem in members.items():
566
+ chosen = select_inner_nodes(mem, states[seq], INNER_MAX_NODES, INNER_MAX_STATES)
567
+ nodes_by_seq[seq] = [[pdb, chain] for _, pdb, chain in chosen]
568
+ rows.extend((pdb, chain, seq, i) for i, (_, pdb, chain) in enumerate(chosen))
569
+ src.executemany("INSERT OR IGNORE INTO sel VALUES (?,?,?,?)", rows)
570
+ log(f" {len(rows):,} sampled observations over {len(nodes_by_seq):,} sequences")
571
+
572
+ edges: dict = defaultdict(list)
573
+ seen: set = set()
574
+ n = 0
575
+ cur = src.execute("""
576
+ SELECT a.seq, a.idx, b.idx, e.pair_similarity, e.pair_fidelity
577
+ FROM temp.sel a
578
+ JOIN edge e ON e.pdb_id_A = a.pdb COLLATE NOCASE AND e.auth_asym_id_A = a.chain
579
+ JOIN temp.sel b ON b.pdb = lower(e.pdb_id_B) AND b.chain = e.auth_asym_id_B AND b.seq = a.seq
580
+ WHERE a.idx < b.idx
581
+ """)
582
+ while True:
583
+ batch = cur.fetchmany(200000)
584
+ if not batch:
585
+ break
586
+ for seq, i, j, sim, fid in batch:
587
+ if (seq, i, j) in seen: # the edge table holds some exact duplicate rows
588
+ continue
589
+ seen.add((seq, i, j))
590
+ simf = to_float(sim)
591
+ edges[seq].append([i, j, -1 if simf is None else round(simf * 1000), FIDELITY_CODE.get(fid, 4)])
592
+ n += len(batch)
593
+ log(f" inner edges: {n:,}")
594
+
595
+ dst.executescript(
596
+ """
597
+ DROP TABLE IF EXISTS inner_graph;
598
+ CREATE TABLE inner_graph (sequence_id TEXT PRIMARY KEY, max_nodes INTEGER, max_states INTEGER,
599
+ nodes TEXT, edges BLOB);
600
+ """
601
+ )
602
+ dst.executemany(
603
+ "INSERT INTO inner_graph VALUES (?,?,?,?,?)",
604
+ (
605
+ (seq, INNER_MAX_NODES, INNER_MAX_STATES, json.dumps(nodes),
606
+ zlib.compress(json.dumps(edges.get(seq, []), separators=(",", ":")).encode(), 6))
607
+ for seq, nodes in nodes_by_seq.items()
608
+ ),
609
+ )
610
+
611
+
612
  # ─────────────────────────────────────────────────────────────── overview
613
  def outer_graph_stats(dst: sqlite3.Connection, n_sequences: int) -> dict:
614
  rows = dst.execute(
 
704
  ap.add_argument("--out-dir", type=Path, required=True)
705
  ap.add_argument("--limit", type=int, default=None, help="debug: only read the first N rows")
706
  ap.add_argument("--ecod", type=Path, default=None, help="ecod.latest.domains.txt (for ECOD names)")
707
+ ap.add_argument("--inner-only", action="store_true",
708
+ help="only (re)build the inner_graph table in an existing MuSProt-index.db")
709
  args = ap.parse_args()
710
 
711
  args.out_dir.mkdir(parents=True, exist_ok=True)
 
716
  src = sqlite3.connect(f"file:{args.db}?mode=ro", uri=True)
717
  src.execute("PRAGMA temp_store = FILE")
718
  src.execute("PRAGMA cache_size = -2000000")
719
+
720
+ if args.inner_only:
721
+ dst = sqlite3.connect(index_path)
722
+ log("inner graphs")
723
+ build_inner_graphs(src, dst)
724
+ dst.commit()
725
+ dst.execute("VACUUM")
726
+ dst.close()
727
+ log(f"done β†’ {index_path} ({index_path.stat().st_size / 1e6:.1f} MB)")
728
+ return
729
+
730
  dst = sqlite3.connect(tmp_path)
731
  dst.execute("PRAGMA journal_mode = OFF")
732
  dst.execute("PRAGMA synchronous = OFF")
 
747
  edges = edge_histograms(src, args.limit)
748
  log("pass 4/4: order / disorder labels")
749
  order = order_stats(src, args.limit)
750
+ log("inner graphs")
751
+ build_inner_graphs(src, dst)
752
 
753
  overview = build_overview(seqs, ov, n_nodes, edges, state_fid)
754
  overview["outer_graph"] = outer_graph_stats(dst, len(seqs))
frontend/src/components/KnowledgeGraph.tsx CHANGED
@@ -1,8 +1,8 @@
1
- import { useMemo, useState } from 'react';
2
  import { Link, useNavigate } from 'react-router-dom';
3
  import { ArrowRight } from 'lucide-react';
4
  import { api, type InnerGraph, type OuterGraph } from '../lib/api';
5
- import { useApi } from '../lib/hooks';
6
  import { FIDELITY_INFO, fidelityColor } from '../lib/fidelity';
7
  import { chainLabel, chainPath, cleanFunction, fmtInt, fmtNum, methodShort, sequencePath } from '../lib/format';
8
  import { ForceGraph, type FGGroup, type FGLink, type FGNode } from './ForceGraph';
@@ -234,6 +234,20 @@ export function KnowledgeGraph({ featured }: { featured: Array<{ id: string; nam
234
  const inner = useApi(`inner:${selected}`, (sig) => api.innerGraph(selected, 60, sig));
235
  const group = outer.data?.groups[0];
236
  const g = inner.data;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
237
  const selectedSeq = group?.sequences.find((s) => s.sequence_id === selected) ?? g?.sequence;
238
 
239
  return (
 
1
+ import { useEffect, useMemo, useState } from 'react';
2
  import { Link, useNavigate } from 'react-router-dom';
3
  import { ArrowRight } from 'lucide-react';
4
  import { api, type InnerGraph, type OuterGraph } from '../lib/api';
5
+ import { prefetch, useApi } from '../lib/hooks';
6
  import { FIDELITY_INFO, fidelityColor } from '../lib/fidelity';
7
  import { chainLabel, chainPath, cleanFunction, fmtInt, fmtNum, methodShort, sequencePath } from '../lib/format';
8
  import { ForceGraph, type FGGroup, type FGLink, type FGNode } from './ForceGraph';
 
234
  const inner = useApi(`inner:${selected}`, (sig) => api.innerGraph(selected, 60, sig));
235
  const group = outer.data?.groups[0];
236
  const g = inner.data;
237
+
238
+ // once the first graph is up, warm the other examples so switching is instant
239
+ const firstReady = !!inner.data;
240
+ useEffect(() => {
241
+ if (!firstReady) return;
242
+ const t = setTimeout(() => {
243
+ for (const f of featured) {
244
+ prefetch(`inner:${f.id}`, () => api.innerGraph(f.id, 60));
245
+ prefetch(`outer1:${f.id}`, () => api.outerGraph([f.id]));
246
+ }
247
+ }, 800);
248
+ return () => clearTimeout(t);
249
+ // eslint-disable-next-line react-hooks/exhaustive-deps
250
+ }, [firstReady]);
251
  const selectedSeq = group?.sequences.find((s) => s.sequence_id === selected) ?? g?.sequence;
252
 
253
  return (
frontend/src/lib/hooks.ts CHANGED
@@ -49,6 +49,16 @@ export function useApi<T>(key: string | null, fetcher: (signal: AbortSignal) =>
49
  return state;
50
  }
51
 
 
 
 
 
 
 
 
 
 
 
52
  /** Resolve a best-effort promise (e.g. an RCSB title) into state. */
53
  export function usePromise<T>(factory: () => Promise<T> | null, deps: unknown[]): T | null {
54
  const [value, setValue] = useState<T | null>(null);
 
49
  return state;
50
  }
51
 
52
+ /** Warm the useApi cache for a key (e.g. the next thing a user is likely to open). */
53
+ export function prefetch<T>(key: string, fetcher: () => Promise<T>): void {
54
+ if (cache.has(key)) return;
55
+ fetcher()
56
+ .then((data) => cache.set(key, data))
57
+ .catch(() => {
58
+ /* best effort */
59
+ });
60
+ }
61
+
62
  /** Resolve a best-effort promise (e.g. an RCSB title) into state. */
63
  export function usePromise<T>(factory: () => Promise<T> | null, deps: unknown[]): T | null {
64
  const [value, setValue] = useState<T | null>(null);