MuSProt / backend /scripts /build_index.py
WinslowFan's picture
Claude Opus 5.5
Sortable observations table, hoverable alternative in the 3D comparison, EM label fallback
047712c
Raw History Blame Contribute Delete
38.5 kB
#!/usr/bin/env python3
"""
Build the website sidecars from MuSProt.db.
Outputs (into --out-dir):
MuSProt-index.db compact, indexed lookup tables the website queries directly
sequence one row per sequence_id (catalog: counts, representative, function)
member one row per chain observation (light columns only), incl. its
biological binding partners (JSON list of partner keys)
partner_name partner key -> type, description, UniProt (from the partner table)
EM entries carry an em_label (cryo-EM, negative-stain EM, MicroED, ...) from the
em_specimen table; sequence.methods and the overview count that label instead
of the bare PDB method, so negative-stain maps are not reported as cryo-EM.
state one row per (sequence_id, state_id)
state_pair mean similarity / fidelity between two states of a sequence
function_fts FTS5 over UniProt / CATH / ECOD / Pfam / top ranked functions per sequence
ecod_h, ecod_f ECOD homology / family names (from ecod.latest.domains.txt)
ecod_group outer-graph groups: sequences sharing the same set of ECOD H-groups
overview.json dataset-level counts and histograms for the home page
ECOD names (homology / family) come from ECOD's own domain list via --ecod.
The main DB is only read. Everything is derived in a few SQL scans, so the
edge table (100M+ rows) is aggregated inside SQLite rather than in Python.
Usage:
python backend/scripts/build_index.py /path/MuSProt.db --out-dir /tmp/musprot-bucket \
--ecod /path/ecod.latest.domains.txt
"""
from __future__ import annotations
import argparse
import json
import math
import sqlite3
import time
from collections import Counter, defaultdict
from pathlib import Path
FUNCTION_TEXT_TOP_N = 3
FUNCTION_TEXT_MAX_CHARS = 600
MAX_PAIR_STATES = 40 # state_pair rows are kept only among a sequence's largest states
INNER_MAX_NODES = 60 # precomputed inner graph: observations per sequence …
INNER_MAX_STATES = 7 # … drawn from its largest states (must match app/protein/catalog.py)
FIDELITY_CODE = {"identical": 0, "low": 1, "medium": 2, "high": 3}
def log(msg: str) -> None:
print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True)
def to_float(v):
try:
f = float(v)
return None if math.isnan(f) else f
except (TypeError, ValueError):
return None
def to_int(v):
f = to_float(v)
return None if f is None else int(f)
def parse_list(raw):
raw = (raw or "").strip()
if not raw:
return []
try:
val = json.loads(raw)
return val if isinstance(val, list) else []
except ValueError:
import ast
try:
val = ast.literal_eval(raw)
return val if isinstance(val, list) else []
except (ValueError, SyntaxError):
return []
def hist(values, edges):
"""Counts of values in [edges[i], edges[i+1]); last bin is closed."""
counts = [0] * (len(edges) - 1)
for v in values:
if v is None:
continue
for i in range(len(edges) - 1):
if v < edges[i + 1] or i == len(edges) - 2:
if v >= edges[i]:
counts[i] += 1
break
return counts
# ─────────────────────────────────────────────────────────────── domain inputs
def split_ids(raw) -> list[str]:
return [x for x in (raw or "").split(";") if x]
def ecod_hset(fids: str) -> tuple:
"""ECOD F-ids ('292.2.1.1;11.1.1.3') -> sorted homology-level set ('11.1', '292.2')."""
return tuple(sorted({".".join(f.split(".")[:2]) for f in split_ids(fids)}))
def load_ecod_names(path: Path | None) -> tuple[dict, dict]:
"""f_id -> (t_name, f_name) and h_id -> (architecture, x_name, h_name)."""
fam, hom = {}, {}
if path is None:
return fam, hom
header = None
with open(path, encoding="utf-8") as fh:
for line in fh:
if line.startswith("#"):
cols = line.lstrip("#").rstrip("\n").split("\t")
if "f_id" in cols:
header = cols
continue
parts = line.rstrip("\n").split("\t")
if header is None:
if "f_id" in parts:
header = parts
continue
r = dict(zip(header, parts))
f_id = r.get("f_id", "")
if not f_id or f_id in fam:
continue
clean = lambda v: (v or "").strip('"') if (v or "").strip('"') not in ("NO_X_NAME", "NO_H_NAME", "NO_T_NAME", "F_UNCLASSIFIED") else ""
fam[f_id] = (clean(r.get("t_name")), clean(r.get("f_name")))
h_id = ".".join(f_id.split(".")[:2])
if h_id not in hom:
hom[h_id] = (clean(r.get("architecture_name")), clean(r.get("x_name")), clean(r.get("h_name")))
return fam, hom
# ─────────────────────────────────────────────────────────────── binding partners
# partner table (pipeline step 12): one row per (target chain, partner chain).
# A contact counts as a biological partner when the target side has at least
# PARTNER_MIN_RES interface residues and the contact exists in biological
# assembly 1 (entries without assembly records: not only via a symmetry mate).
# Must match app/protein/records.py.
PARTNER_MIN_RES = 5
NUCLEIC_LABEL = {"dna": "DNA", "rna": "RNA", "hybrid": "DNA/RNA hybrid"}
def partner_is_biological(n_res_target, via_symmetry, in_assembly1) -> bool:
if (to_int(n_res_target) or 0) < PARTNER_MIN_RES:
return False
if in_assembly1 in ("0", "1"):
return in_assembly1 == "1"
return via_symmetry == "0"
def _sequence_like(desc: str) -> bool:
"""Nucleic-acid descriptions that just spell the sequence, e.g. DNA (5'-D(*CP*GP*...)-3')."""
d = desc.strip().upper()
return "*" in d or d.startswith(("5'", "DNA (", "RNA (", "DNA(", "RNA(")) or d in ("DNA", "RNA")
def partner_key(ptype: str, desc: str, uniprot: str) -> str:
"""UniProt accession(s) when known; otherwise '<type>:<description>', or just
the polymer type for nucleic acids whose description only spells the sequence."""
if uniprot:
return uniprot
if not desc or (ptype in NUCLEIC_LABEL and _sequence_like(desc)):
return ptype or "other"
d = " ".join(desc.lower().split()).replace("ribosomal rna", "rrna")
return f"{ptype}:{d}"
def load_partners(src: sqlite3.Connection):
"""(pdb_lower, chain) -> [(key, n_res_target)] biological partners, largest
interface first; 'self' marks another copy of the same entity. Also returns
key -> Counter of (type, description, uniprot) for naming."""
has = src.execute("SELECT 1 FROM sqlite_master WHERE type='table' AND name='partner'").fetchone()
if not has:
log(" no partner table in the DB; skipping binding partners")
return {}, {}
per_chain: dict = defaultdict(dict)
names: dict = defaultdict(Counter)
n = 0
for pdb, chain, ptype, desc, unp, nres, same, sym, asm in src.execute(
"SELECT pdb_id, auth_asym_id, partner_type, partner_description, partner_uniprot, "
"n_res_target, same_entity, via_symmetry, in_assembly1 FROM partner"
):
n += 1
if not partner_is_biological(nres, sym, asm):
continue
if same == "1":
key = "self"
else:
key = partner_key(ptype or "", desc or "", unp or "")
names[key][(ptype or "", desc or "", unp or "")] += 1
d = per_chain[(pdb.lower(), chain)]
d[key] = max(d.get(key, 0), to_int(nres) or 0)
log(f" partner rows: {n:,}; chains with a biological partner: {len(per_chain):,}")
out = {k: sorted(v.items(), key=lambda kv: -kv[1]) for k, v in per_chain.items()}
return out, names
def write_partner_names(dst: sqlite3.Connection, names: dict, used: Counter) -> None:
dst.executescript(
"""
DROP TABLE IF EXISTS partner_name;
CREATE TABLE partner_name (
key TEXT PRIMARY KEY,
type TEXT,
description TEXT,
uniprot TEXT,
n_chains INTEGER
);
"""
)
rows = []
for key, n_chains in used.items():
if key == "self":
continue
(ptype, desc, unp), _ = names[key].most_common(1)[0]
if key in NUCLEIC_LABEL:
desc = NUCLEIC_LABEL[key]
rows.append((key, ptype, desc, unp or None, n_chains))
dst.executemany("INSERT INTO partner_name VALUES (?,?,?,?,?)", rows)
# ─────────────────────────────────────────────────────────────── EM specimen
def load_em_labels(src: sqlite3.Connection) -> dict:
"""pdb_lower -> em_label: from the em_specimen table (step 13 patch) or, in a
DB rebuilt from the merged step 1, from node.em_label; {} when neither exists."""
has = src.execute("SELECT 1 FROM sqlite_master WHERE type='table' AND name='em_specimen'").fetchone()
if has:
return {p.lower(): lab for p, lab in src.execute("SELECT pdb_id, em_label FROM em_specimen") if lab}
if "em_label" in {r[1] for r in src.execute("PRAGMA table_info(node)")}:
return {p.lower(): lab for p, lab in src.execute(
"SELECT DISTINCT pdb_id, em_label FROM node WHERE em_label <> ''") if lab}
log(" no EM specimen labels in the DB; EM entries keep the bare PDB method")
return {}
# ─────────────────────────────────────────────────────────────── node pass
NODE_SQL = """
SELECT sequence_id, uniprot_id, pdb_id, auth_asym_id, sequence_length,
n_resolved_aa, resolved_coverage, binders, binding_status,
experimental_method, resolution, pH, temp_K, chain_composition,
non_protein_polymer_binding, initial_release_date, cath_superfamily,
state_id, ranked_functions, cath_id, pfam_id, ecod_id, ecod_fid
FROM node
"""
def build_members(src: sqlite3.Connection, dst: sqlite3.Connection, limit: int | None):
dst.executescript(
"""
DROP TABLE IF EXISTS member;
CREATE TABLE member (
sequence_id TEXT NOT NULL,
state_id INTEGER,
pdb_id TEXT NOT NULL,
auth_asym_id TEXT NOT NULL,
uniprot_id TEXT,
experimental_method TEXT,
resolution REAL,
binding_status TEXT,
chain_composition TEXT,
binders TEXT,
release_date TEXT,
resolved_coverage REAL,
n_resolved_aa INTEGER,
pH REAL,
temp_K REAL,
cath_id TEXT,
pfam_id TEXT,
ecod_id TEXT,
ecod_fid TEXT,
partners TEXT,
em_label TEXT
);
"""
)
partners, partner_names = load_partners(src)
em_labels = load_em_labels(src)
seqs: dict[str, dict] = {}
functions: dict[str, list[str]] = {}
ov = {
"method": Counter(), "binding": Counter(), "composition": Counter(),
"npp": Counter(), "year": Counter(), "cath": Counter(),
"resolution": [], "length": {}, "uniprot": set(), "entries": set(),
"partner_used": Counter(), "partner_type": Counter(),
"with_partner": 0, "with_hetero_partner": 0,
}
sql = NODE_SQL + (f" LIMIT {int(limit)}" if limit else "")
batch = []
n = 0
for row in src.execute(sql):
(seq_id, uniprot, pdb, chain, seq_len, n_res, res_cov, binders, binding,
method, resolution, ph, temp, comp, npp, date, cath_sf, state, funcs, cath_id,
pfam, ecod_ids, ecod_fid) = row
state_i = to_int(state)
res_f = to_float(resolution)
em_label = em_labels.get(pdb.lower())
method_label = em_label or method # what the site shows and counts
# another entity with the same UniProt accession is still a homo contact
pkeys = []
for key, _ in partners.get((pdb.lower(), chain), ()):
key = "self" if uniprot and key == uniprot else key
if key not in pkeys:
pkeys.append(key)
if pkeys:
ov["with_partner"] += 1
hetero = [k for k in pkeys if k != "self"]
ov["with_hetero_partner"] += bool(hetero)
ov["partner_used"].update(pkeys)
ov["partner_type"].update({partner_names[k].most_common(1)[0][0][0] for k in hetero})
batch.append((
seq_id, state_i, pdb, chain, uniprot or None, method or None, res_f,
binding or None, comp or None, binders or None, date or None,
to_float(res_cov), to_int(n_res), to_float(ph), to_float(temp),
cath_id or None, pfam or None, ecod_ids or None, ecod_fid or None,
json.dumps(pkeys, separators=(",", ":")) if pkeys else None,
em_label,
))
s = seqs.get(seq_id)
if s is None:
s = seqs[seq_id] = {
"uniprot": Counter(), "length": to_int(seq_len), "n_obs": 0,
"states": Counter(), "n_holo": 0, "entries": set(), "methods": Counter(),
"cath": Counter(), "dates": [], "rep": None,
"ecod": Counter(), "pfam": Counter(),
}
s["n_obs"] += 1
if uniprot:
s["uniprot"][uniprot] += 1
s["states"][state_i] += 1
s["n_holo"] += binding == "holo"
s["entries"].add(pdb)
if method_label:
s["methods"][method_label] += 1
if cath_sf:
s["cath"][cath_sf] += 1
hs = ecod_hset(ecod_fid)
if hs:
s["ecod"][hs] += 1
pf = tuple(sorted(set(split_ids(pfam))))
if pf:
s["pfam"][pf] += 1
if date:
s["dates"].append(date)
# representative: best (lowest) resolution among well-resolved chains
key = ((to_float(res_cov) or 0) < 0.9, res_f if res_f is not None else 99.0)
if s["rep"] is None or key < s["rep"][0]:
s["rep"] = (key, pdb, chain)
if seq_id not in functions and funcs:
fl = [str(f).strip() for f in parse_list(funcs) if len(str(f).strip()) > 10]
if fl:
functions[seq_id] = fl[:FUNCTION_TEXT_TOP_N]
ov["method"][method_label or "Unknown"] += 1
ov["binding"][binding or "Unknown"] += 1
ov["composition"][comp or "Unknown"] += 1
ov["npp"][npp or "None"] += 1
if date:
ov["year"][date[:4]] += 1
if res_f is not None:
ov["resolution"].append(res_f)
if uniprot:
ov["uniprot"].add(uniprot)
ov["entries"].add(pdb.lower())
ov["length"][seq_id] = to_int(seq_len)
n += 1
if len(batch) >= 50000:
dst.executemany(f"INSERT INTO member VALUES ({','.join('?' * 21)})", batch)
batch.clear()
log(f" node rows: {n:,}")
if batch:
dst.executemany(f"INSERT INTO member VALUES ({','.join('?' * 21)})", batch)
log(f" node rows total: {n:,}; sequences: {len(seqs):,}")
write_partner_names(dst, partner_names, ov["partner_used"])
top = []
for key, c in ov["partner_used"].most_common(41):
if key == "self":
continue
(ptype, desc, unp), _ = partner_names[key].most_common(1)[0]
top.append({"key": key, "type": ptype, "label": NUCLEIC_LABEL.get(key, desc),
"uniprot": unp or None, "value": c})
ov["partners"] = {
"observations_with_partner": ov["with_partner"],
"observations_with_hetero_partner": ov["with_hetero_partner"],
"observations_homo_only": ov["with_partner"] - ov["with_hetero_partner"],
"distinct_partners": len([k for k in ov["partner_used"] if k != "self"]),
"by_type": [{"label": k, "value": v} for k, v in ov["partner_type"].most_common()],
"top": top[:40],
}
dst.executescript(
"""
CREATE INDEX idx_member_seq ON member(sequence_id, state_id);
CREATE INDEX idx_member_pdb ON member(pdb_id COLLATE NOCASE, auth_asym_id);
CREATE INDEX idx_member_uniprot ON member(uniprot_id COLLATE NOCASE);
"""
)
return seqs, functions, ov, n
def consensus(counter: Counter) -> tuple:
"""Most common non-empty assignment set across a sequence's chains (ties: larger set)."""
if not counter:
return ()
return max(counter.items(), key=lambda kv: (kv[1], len(kv[0]), kv[0]))[0]
def write_sequences(dst: sqlite3.Connection, seqs: dict, functions: dict, ecod_names: tuple):
dst.executescript(
"""
DROP TABLE IF EXISTS sequence;
CREATE TABLE sequence (
sequence_id TEXT PRIMARY KEY,
uniprot_id TEXT,
length INTEGER,
n_obs INTEGER,
n_states INTEGER,
n_entries INTEGER,
n_holo INTEGER,
n_apo INTEGER,
n_cross_pairs INTEGER,
largest_state INTEGER,
rep_pdb TEXT,
rep_chain TEXT,
cath_superfamily TEXT,
methods TEXT,
first_release TEXT,
last_release TEXT,
top_function TEXT,
ecod_hset TEXT,
pfam_ids TEXT
);
DROP TABLE IF EXISTS state;
CREATE TABLE state (
sequence_id TEXT NOT NULL,
state_id INTEGER NOT NULL,
n_members INTEGER,
PRIMARY KEY (sequence_id, state_id)
);
"""
)
rows, state_rows = [], []
for seq_id, s in seqs.items():
sizes = list(s["states"].values())
n_obs = s["n_obs"]
cross = n_obs * n_obs - sum(x * x for x in sizes) # directed cross-state pairs
dates = sorted(s["dates"])
rows.append((
seq_id,
s["uniprot"].most_common(1)[0][0] if s["uniprot"] else None,
s["length"], n_obs, len(sizes), len(s["entries"]), s["n_holo"], n_obs - s["n_holo"],
cross, max(sizes),
s["rep"][1], s["rep"][2],
s["cath"].most_common(1)[0][0] if s["cath"] else None,
";".join(m for m, _ in s["methods"].most_common()),
dates[0] if dates else None, dates[-1] if dates else None,
(functions.get(seq_id) or [None])[0],
";".join(consensus(s["ecod"])) or None,
";".join(consensus(s["pfam"])) or None,
))
for st, cnt in s["states"].items():
state_rows.append((seq_id, st, cnt))
dst.executemany(f"INSERT INTO sequence VALUES ({','.join('?' * 19)})", rows)
dst.executemany("INSERT INTO state VALUES (?,?,?)", state_rows)
dst.executescript(
"""
CREATE INDEX idx_seq_states ON sequence(n_states DESC, n_obs DESC);
CREATE INDEX idx_seq_obs ON sequence(n_obs DESC);
CREATE INDEX idx_seq_uniprot ON sequence(uniprot_id COLLATE NOCASE);
CREATE INDEX idx_seq_cath ON sequence(cath_superfamily);
CREATE INDEX idx_seq_ecod ON sequence(ecod_hset, n_obs DESC);
"""
)
fam, hom = ecod_names
dst.executescript(
"""
DROP TABLE IF EXISTS ecod_h;
CREATE TABLE ecod_h (h_id TEXT PRIMARY KEY, architecture TEXT, x_name TEXT, h_name TEXT);
DROP TABLE IF EXISTS ecod_f;
CREATE TABLE ecod_f (f_id TEXT PRIMARY KEY, t_name TEXT, f_name TEXT);
DROP TABLE IF EXISTS ecod_group;
CREATE TABLE ecod_group (
hset TEXT PRIMARY KEY,
label TEXT,
n_sequences INTEGER,
n_obs INTEGER,
n_multistate INTEGER,
n_states INTEGER
);
"""
)
dst.executemany("INSERT INTO ecod_h VALUES (?,?,?,?)", [(k, *v) for k, v in hom.items()])
dst.executemany("INSERT INTO ecod_f VALUES (?,?,?)", [(k, *v) for k, v in fam.items()])
groups: dict[tuple, list] = {}
for s in seqs.values():
hs = consensus(s["ecod"])
if hs:
g = groups.setdefault(hs, [0, 0, 0, 0])
g[0] += 1
g[1] += s["n_obs"]
g[2] += len(s["states"]) > 1
g[3] += len(s["states"])
dst.executemany(
"INSERT INTO ecod_group VALUES (?,?,?,?,?,?)",
[
(";".join(hs), " + ".join(hom.get(h, ("", "", ""))[2] or h for h in hs), *v)
for hs, v in groups.items()
],
)
dst.execute("CREATE INDEX idx_group_size ON ecod_group(n_sequences DESC)")
dst.executescript(
"""
DROP TABLE IF EXISTS function_fts;
CREATE VIRTUAL TABLE function_fts USING fts5(
sequence_id UNINDEXED, uniprot_id, cath_superfamily, domains, text,
tokenize = 'porter unicode61'
);
"""
)
fts_rows = []
for seq_id, s in seqs.items():
text = " ".join(functions.get(seq_id, []))[:FUNCTION_TEXT_MAX_CHARS * FUNCTION_TEXT_TOP_N]
uni = " ".join(s["uniprot"].keys())
cath = " ".join(s["cath"].keys()).replace(";", " ")
hs = consensus(s["ecod"])
dom = " ".join(
[*(h.replace(".", "_") for h in hs), *(hom.get(h, ("", "", ""))[2] for h in hs),
*consensus(s["pfam"])]
)
fts_rows.append((seq_id, uni, cath, dom, text))
dst.executemany("INSERT INTO function_fts VALUES (?,?,?,?,?)", fts_rows)
# ─────────────────────────────────────────────────────────────── edge passes
def build_state_pairs(src: sqlite3.Connection, dst: sqlite3.Connection, seqs: dict,
limit: int | None) -> dict:
"""Store state-pair similarities (top states only); return fidelity counts over all pairs."""
keep = {
seq_id: {st for st, _ in s["states"].most_common(MAX_PAIR_STATES)}
for seq_id, s in seqs.items() if len(s["states"]) > MAX_PAIR_STATES
}
fid_counts: dict = defaultdict(lambda: {"state_pairs": 0, "observation_pairs": 0})
dst.executescript(
"""
DROP TABLE IF EXISTS state_pair;
CREATE TABLE state_pair (
sequence_id TEXT NOT NULL,
state_a INTEGER NOT NULL,
state_b INTEGER NOT NULL,
similarity REAL,
fidelity TEXT,
n_pairs INTEGER,
PRIMARY KEY (sequence_id, state_a, state_b)
);
"""
)
src_tbl = f"(SELECT * FROM edge LIMIT {int(limit)})" if limit else "edge"
# directed edges -> keep one direction (A <= B) and count undirected pairs
sql = f"""
SELECT sequence_id, CAST(state_id_A AS INTEGER) a, CAST(state_id_B AS INTEGER) b,
MIN(CAST(state_similarity AS REAL)), MIN(state_fidelity), COUNT(*)
FROM {src_tbl}
WHERE CAST(state_id_A AS INTEGER) <= CAST(state_id_B AS INTEGER)
GROUP BY sequence_id, a, b
"""
cur = src.execute(sql)
n = 0
while True:
rows = cur.fetchmany(100000)
if not rows:
break
out = []
for s, a, b, sim, fid, c in rows:
# within-state rows appear in both directions with A==B -> halve the count
c = c // 2 if a == b else c
if a != b:
fc = fid_counts[fid or "NA"]
fc["state_pairs"] += 1
fc["observation_pairs"] += c
allowed = keep.get(s)
if allowed is None or (a in allowed and b in allowed):
out.append((s, a, b, sim, fid, c))
dst.executemany("INSERT INTO state_pair VALUES (?,?,?,?,?,?)", out)
n += len(rows)
log(f" state pairs: {n:,}")
return dict(fid_counts)
def edge_histograms(src: sqlite3.Connection, limit: int | None) -> dict:
src_tbl = f"(SELECT * FROM edge LIMIT {int(limit)})" if limit else "edge"
sql = f"""
SELECT state_id_A = state_id_B AS same_state,
COALESCE(NULLIF(pair_fidelity, ''), 'NA') AS fid,
-- pairs without shared residues store '' similarity and 0.0 RMSD/coverage: keep them
-- out of every histogram (they are counted under fidelity 'NA')
CASE WHEN pair_similarity = '' THEN NULL
ELSE MIN(CAST(CAST(pair_similarity AS REAL) * 50 AS INTEGER), 49) END AS sim_bin,
CASE WHEN pair_similarity = '' THEN NULL
ELSE MIN(CAST(CAST(RMSD AS REAL) * 2 AS INTEGER), 40) END AS rmsd_bin,
CASE WHEN pair_similarity = '' THEN NULL
ELSE MIN(CAST(CAST(coverage AS REAL) * 20 AS INTEGER), 19) END AS cov_bin,
CASE WHEN tm_aln = '' THEN NULL
ELSE MIN(CAST(CAST(tm_aln AS REAL) * 50 AS INTEGER), 49) END AS tm_bin,
COUNT(*)
FROM {src_tbl}
GROUP BY 1, 2, 3, 4, 5, 6
"""
out = {
"total": 0, "cross_state": 0,
"pair_fidelity": Counter(), "pair_fidelity_cross": Counter(),
"similarity": [0] * 50, "similarity_cross": [0] * 50,
"tm_aln": [0] * 50, "rmsd": [0] * 41, "coverage": [0] * 20,
}
for same, fid, sim_b, rmsd_b, cov_b, tm_b, c in src.execute(sql):
out["total"] += c
out["pair_fidelity"][fid] += c
if not same:
out["cross_state"] += c
out["pair_fidelity_cross"][fid] += c
if sim_b is not None and sim_b >= 0:
out["similarity"][sim_b] += c
if not same:
out["similarity_cross"][sim_b] += c
if tm_b is not None and tm_b >= 0:
out["tm_aln"][tm_b] += c
if rmsd_b is not None and rmsd_b >= 0:
out["rmsd"][rmsd_b] += c
if cov_b is not None and cov_b >= 0:
out["coverage"][cov_b] += c
return out
def order_stats(src: sqlite3.Connection, limit: int | None) -> dict:
"""Counts of the observed order/disorder labels over all directed transitions."""
src_tbl = f"(SELECT * FROM edge LIMIT {int(limit)})" if limit else "edge"
rows = src.execute(f"""
SELECT order_invalid_reason, order_evidence, ordering_with_ligand_change,
state_id_A = state_id_B, COUNT(*)
FROM {src_tbl} GROUP BY 1, 2, 3, 4
""").fetchall()
reason, evidence, evidence_cross, ligand = Counter(), Counter(), Counter(), Counter()
for r, ev, lig, same, c in rows:
reason[r or "NA"] += c
if r == "valid":
evidence[ev or "NA"] += c
if not same:
evidence_cross[ev or "NA"] += c
if ev in ("low", "medium", "high"):
ligand[(ev, lig == "True")] += c
order = ["high", "medium", "low", "not_applicable"]
return {
"validity": [{"label": k, "value": v} for k, v in reason.most_common()],
"evidence": [{"label": k, "value": evidence.get(k, 0), "cross_state": evidence_cross.get(k, 0)} for k in order],
"with_ligand_change": [
{"label": k, "with": ligand.get((k, True), 0), "without": ligand.get((k, False), 0)}
for k in order[:3]
],
}
def select_inner_nodes(members: list, states: list, max_nodes: int, max_states: int) -> list:
"""Round-robin over the largest states. Mirrors catalog.inner_graph_nodes exactly."""
order = [st for st, _ in states[:max_states]]
keep = set(order)
pool = [m for m in members if m[0] in keep]
if len(pool) <= max_nodes:
return pool
by_state: dict = {}
for m in pool:
by_state.setdefault(m[0], []).append(m)
chosen, depth = [], 0
while len(chosen) < max_nodes:
added = False
for st in order:
bucket = by_state.get(st, [])
if depth < len(bucket) and len(chosen) < max_nodes:
chosen.append(bucket[depth])
added = True
if not added:
break
depth += 1
return chosen
def build_inner_graphs(src: sqlite3.Connection, dst: sqlite3.Connection) -> None:
"""Precompute each sequence's inner-graph sample and its pairwise similarities.
Serving these live needs one edge-index lookup per node, which is slow when the
DB sits on a network mount; stored here it is a single row read.
"""
import zlib
states: dict = defaultdict(list)
for seq, st, n in dst.execute(
"SELECT sequence_id, state_id, n_members FROM state ORDER BY sequence_id, n_members DESC, state_id"
):
states[seq].append((st, n))
members: dict = defaultdict(list)
for seq, st, pdb, chain in dst.execute(
"SELECT sequence_id, state_id, pdb_id, auth_asym_id FROM member "
"ORDER BY sequence_id, state_id, resolution IS NULL, resolution"
):
members[seq].append((st, pdb.lower(), chain))
src.execute("CREATE TEMP TABLE sel (pdb TEXT, chain TEXT, seq TEXT, idx INTEGER, PRIMARY KEY (pdb, chain))")
nodes_by_seq = {}
rows = []
for seq, mem in members.items():
chosen = select_inner_nodes(mem, states[seq], INNER_MAX_NODES, INNER_MAX_STATES)
nodes_by_seq[seq] = [[pdb, chain] for _, pdb, chain in chosen]
rows.extend((pdb, chain, seq, i) for i, (_, pdb, chain) in enumerate(chosen))
src.executemany("INSERT OR IGNORE INTO sel VALUES (?,?,?,?)", rows)
log(f" {len(rows):,} sampled observations over {len(nodes_by_seq):,} sequences")
edges: dict = defaultdict(list)
seen: set = set()
n = 0
cur = src.execute("""
SELECT a.seq, a.idx, b.idx, e.pair_similarity, e.pair_fidelity
FROM temp.sel a
JOIN edge e ON e.pdb_id_A = a.pdb COLLATE NOCASE AND e.auth_asym_id_A = a.chain
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
WHERE a.idx < b.idx
""")
while True:
batch = cur.fetchmany(200000)
if not batch:
break
for seq, i, j, sim, fid in batch:
if (seq, i, j) in seen: # the edge table holds some exact duplicate rows
continue
seen.add((seq, i, j))
simf = to_float(sim)
edges[seq].append([i, j, -1 if simf is None else round(simf * 1000), FIDELITY_CODE.get(fid, 4)])
n += len(batch)
log(f" inner edges: {n:,}")
dst.executescript(
"""
DROP TABLE IF EXISTS inner_graph;
CREATE TABLE inner_graph (sequence_id TEXT PRIMARY KEY, max_nodes INTEGER, max_states INTEGER,
nodes TEXT, edges BLOB);
"""
)
dst.executemany(
"INSERT INTO inner_graph VALUES (?,?,?,?,?)",
(
(seq, INNER_MAX_NODES, INNER_MAX_STATES, json.dumps(nodes),
zlib.compress(json.dumps(edges.get(seq, []), separators=(",", ":")).encode(), 6))
for seq, nodes in nodes_by_seq.items()
),
)
# ─────────────────────────────────────────────────────────────── overview
def outer_graph_stats(dst: sqlite3.Connection, n_sequences: int) -> dict:
rows = dst.execute(
"SELECT hset, label, n_sequences, n_obs, n_multistate FROM ecod_group ORDER BY n_sequences DESC"
).fetchall()
covered = sum(r[2] for r in rows)
size_bins = Counter()
for r in rows:
n = r[2]
size_bins["1" if n == 1 else "2–5" if n <= 5 else "6–20" if n <= 20 else "21–100" if n <= 100 else ">100"] += 1
return {
"level": "ECOD homology (H)",
"groups": len(rows),
"covered_sequences": covered,
"coverage": covered / n_sequences if n_sequences else 0,
"edges": sum(r[2] * (r[2] - 1) // 2 for r in rows),
"group_sizes": [{"label": k, "value": size_bins[k]} for k in ["1", "2–5", "6–20", "21–100", ">100"]],
"largest": [
{"hset": r[0], "label": r[1], "n_sequences": r[2], "n_obs": r[3], "n_multistate": r[4]}
for r in rows[:12]
],
}
def build_overview(seqs, ov, n_nodes, edges, state_fid) -> dict:
states_per_seq = Counter(len(s["states"]) for s in seqs.values())
obs_per_seq = [s["n_obs"] for s in seqs.values()]
n_states_total = sum(len(s["states"]) for s in seqs.values())
multi_state = sum(1 for s in seqs.values() if len(s["states"]) > 1)
lengths = [v for v in ov["length"].values() if v]
def capped(counter, cap):
out = Counter()
for k, v in counter.items():
out[k if k < cap else cap] += v
return [{"label": (f"{k}+" if k == cap else str(k)), "value": out[k]} for k in sorted(out)]
obs_edges = [2, 3, 4, 5, 6, 11, 21, 51, 101, 1_000_000]
obs_labels = ["2", "3", "4", "5", "6–10", "11–20", "21–50", "51–100", ">100"]
len_edges = [0, 100, 200, 300, 400, 500, 750, 1000, 1_000_000]
len_labels = ["<100", "100–199", "200–299", "300–399", "400–499", "500–749", "750–999", "β‰₯1000"]
res_edges = [0, 1.5, 2.0, 2.5, 3.0, 3.5, 4.0, 5.0, 1000]
res_labels = ["<1.5", "1.5–2.0", "2.0–2.5", "2.5–3.0", "3.0–3.5", "3.5–4.0", "4.0–5.0", "β‰₯5.0"]
fid_order = ["identical", "low", "medium", "high", "no_shared_residues", "NA"]
return {
"generated": time.strftime("%Y-%m-%d"),
"counts": {
"observations": n_nodes,
"sequences": len(seqs),
"multi_state_sequences": multi_state,
"state_clusters": n_states_total,
"pdb_entries": len(ov["entries"]),
"uniprot_accessions": len(ov["uniprot"]),
"transitions": edges["total"],
"cross_state_transitions": edges["cross_state"],
},
"states_per_sequence": capped(states_per_seq, 10),
"observations_per_sequence": [
{"label": l, "value": v} for l, v in zip(obs_labels, hist(obs_per_seq, obs_edges))
],
"sequence_length": [
{"label": l, "value": v} for l, v in zip(len_labels, hist(lengths, len_edges))
],
"resolution": [
{"label": l, "value": v} for l, v in zip(res_labels, hist(ov["resolution"], res_edges))
],
"experimental_method": [{"label": k, "value": v} for k, v in ov["method"].most_common()],
"binding_status": [{"label": k, "value": v} for k, v in ov["binding"].most_common()],
"chain_composition": [{"label": k, "value": v} for k, v in ov["composition"].most_common()],
"nucleic_acid_binding": [{"label": k, "value": v} for k, v in ov["npp"].most_common()],
"release_year": [{"label": k, "value": ov["year"][k]} for k in sorted(ov["year"])],
"partners": ov.get("partners"),
"pair_fidelity": [
{"label": k, "value": edges["pair_fidelity"].get(k, 0),
"cross_state": edges["pair_fidelity_cross"].get(k, 0)}
for k in fid_order if edges["pair_fidelity"].get(k)
],
"state_fidelity": [
{"label": k, **state_fid[k]} for k in fid_order if k in state_fid
],
"pair_similarity": {
"bin_width": 0.02, "all": edges["similarity"], "cross_state": edges["similarity_cross"],
},
"tm_aln": {"bin_width": 0.02, "all": edges["tm_aln"]},
"rmsd": {"bin_width": 0.5, "all": edges["rmsd"], "last_bin_open": True},
"coverage": {"bin_width": 0.05, "all": edges["coverage"]},
}
def main():
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("db", type=Path)
ap.add_argument("--out-dir", type=Path, required=True)
ap.add_argument("--limit", type=int, default=None, help="debug: only read the first N rows")
ap.add_argument("--ecod", type=Path, default=None, help="ecod.latest.domains.txt (for ECOD names)")
ap.add_argument("--inner-only", action="store_true",
help="only (re)build the inner_graph table in an existing MuSProt-index.db")
args = ap.parse_args()
args.out_dir.mkdir(parents=True, exist_ok=True)
index_path = args.out_dir / "MuSProt-index.db"
tmp_path = index_path.with_suffix(".db.tmp")
tmp_path.unlink(missing_ok=True)
src = sqlite3.connect(f"file:{args.db}?mode=ro", uri=True)
src.execute("PRAGMA temp_store = FILE")
src.execute("PRAGMA cache_size = -2000000")
if args.inner_only:
dst = sqlite3.connect(index_path)
log("inner graphs")
build_inner_graphs(src, dst)
dst.commit()
dst.execute("VACUUM")
dst.close()
log(f"done β†’ {index_path} ({index_path.stat().st_size / 1e6:.1f} MB)")
return
dst = sqlite3.connect(tmp_path)
dst.execute("PRAGMA journal_mode = OFF")
dst.execute("PRAGMA synchronous = OFF")
ecod_names = load_ecod_names(args.ecod)
log(f"ECOD names: {len(ecod_names[0]):,} families, {len(ecod_names[1]):,} homology groups")
log("pass 1/4: node table")
seqs, functions, ov, n_nodes = build_members(src, dst, args.limit)
write_sequences(dst, seqs, functions, ecod_names)
dst.commit()
log("pass 2/4: state pairs (edge GROUP BY)")
state_fid = build_state_pairs(src, dst, seqs, args.limit)
dst.commit()
log("pass 3/4: edge histograms")
edges = edge_histograms(src, args.limit)
log("pass 4/4: order / disorder labels")
order = order_stats(src, args.limit)
log("inner graphs")
build_inner_graphs(src, dst)
overview = build_overview(seqs, ov, n_nodes, edges, state_fid)
overview["outer_graph"] = outer_graph_stats(dst, len(seqs))
overview["order"] = order
dst.execute("CREATE TABLE meta (key TEXT PRIMARY KEY, value TEXT)")
dst.execute("INSERT INTO meta VALUES ('overview', ?)", (json.dumps(overview),))
dst.commit()
dst.execute("VACUUM")
dst.close()
tmp_path.replace(index_path)
(args.out_dir / "overview.json").write_text(json.dumps(overview, indent=1))
log(f"done β†’ {index_path} ({index_path.stat().st_size / 1e6:.1f} MB)")
log(json.dumps(overview["counts"], indent=1))
if __name__ == "__main__":
main()