Claude Opus 5.5
Sortable observations table, hoverable alternative in the 3D comparison, EM label fallback
047712c Download backend/scripts/build_index.py from omaib/MuSProt: direct link, hf CLI and curl.
- Browser
- Download file 38.5 kB
-
https://huggingface.co/spaces/omaib/MuSProt/resolve/main/backend/scripts/build_index.py
- Command line
-
hf download hf://spaces/omaib/MuSProt/backend/scripts/build_index.py
-
curl -L -o build_index.py https://huggingface.co/spaces/omaib/MuSProt/resolve/main/backend/scripts/build_index.py
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() | |