protcross 0.1.1__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (44) hide show
  1. evopoint_da/__init__.py +4 -0
  2. evopoint_da/assets.py +178 -0
  3. evopoint_da/cli/__init__.py +2 -0
  4. evopoint_da/cli/download_af2.py +37 -0
  5. evopoint_da/cli/main.py +42 -0
  6. evopoint_da/cli/map_labels.py +43 -0
  7. evopoint_da/cli/predict.py +156 -0
  8. evopoint_da/cli/preprocess.py +45 -0
  9. evopoint_da/cli/setup_assets.py +51 -0
  10. evopoint_da/cli/train.py +19 -0
  11. evopoint_da/data/__init__.py +19 -0
  12. evopoint_da/data/af2.py +110 -0
  13. evopoint_da/data/components.py +19 -0
  14. evopoint_da/data/datamodule.py +82 -0
  15. evopoint_da/data/dataset.py +143 -0
  16. evopoint_da/data/esm.py +91 -0
  17. evopoint_da/data/label_mapping.py +340 -0
  18. evopoint_da/data/pca.py +37 -0
  19. evopoint_da/data/preprocess.py +130 -0
  20. evopoint_da/data/structure.py +138 -0
  21. evopoint_da/evaluation/__init__.py +6 -0
  22. evopoint_da/evaluation/adaptive.py +114 -0
  23. evopoint_da/evaluation/metrics.py +78 -0
  24. evopoint_da/experiments/__init__.py +2 -0
  25. evopoint_da/experiments/multiseed_benchmark.py +157 -0
  26. evopoint_da/experiments/strategy_search.py +161 -0
  27. evopoint_da/inference/__init__.py +11 -0
  28. evopoint_da/inference/pdb.py +40 -0
  29. evopoint_da/inference/predictor.py +394 -0
  30. evopoint_da/models/__init__.py +6 -0
  31. evopoint_da/models/backbones/__init__.py +6 -0
  32. evopoint_da/models/backbones/pointnet2.py +216 -0
  33. evopoint_da/models/domain_weights.py +50 -0
  34. evopoint_da/models/heads/__init__.py +6 -0
  35. evopoint_da/models/heads/classifier.py +26 -0
  36. evopoint_da/models/module.py +181 -0
  37. evopoint_da/training/__init__.py +6 -0
  38. evopoint_da/training/run.py +62 -0
  39. protcross-0.1.1.dist-info/METADATA +561 -0
  40. protcross-0.1.1.dist-info/RECORD +44 -0
  41. protcross-0.1.1.dist-info/WHEEL +5 -0
  42. protcross-0.1.1.dist-info/entry_points.txt +8 -0
  43. protcross-0.1.1.dist-info/licenses/LICENSE +21 -0
  44. protcross-0.1.1.dist-info/top_level.txt +1 -0
@@ -0,0 +1,4 @@
1
+ """ProtCross core package."""
2
+
3
+ __version__ = "0.1.1"
4
+
evopoint_da/assets.py ADDED
@@ -0,0 +1,178 @@
1
+ """Download and validate ProtCross runtime assets."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hashlib
6
+ import os
7
+ import shlex
8
+ from dataclasses import dataclass
9
+ from pathlib import Path
10
+
11
+ import requests
12
+ from tqdm import tqdm
13
+
14
+
15
+ RELEASE_TAG = "v0.1.1"
16
+ GITHUB_RELEASE_BASE = f"https://github.com/GeraltZeroZhong/ProtCross/releases/download/{RELEASE_TAG}"
17
+ DEFAULT_ASSETS_DIR = Path.home() / ".cache" / "protcross" / "assets" / RELEASE_TAG
18
+
19
+ DEFAULT_ESM_URL = (
20
+ "https://huggingface.co/EvolutionaryScale/esmc-600m-2024-12/"
21
+ "resolve/main/data/weights/esmc_600m_2024_12_v0.pth"
22
+ )
23
+ DEFAULT_CHECKPOINT_URL = f"{GITHUB_RELEASE_BASE}/best-epoch.59.ckpt"
24
+ DEFAULT_PCA_URL = f"{GITHUB_RELEASE_BASE}/pca_esmc_128.pkl"
25
+
26
+
27
+ @dataclass(frozen=True)
28
+ class AssetSpec:
29
+ name: str
30
+ filename: str
31
+ url: str
32
+ sha256: str | None = None
33
+
34
+
35
+ DEFAULT_ASSETS = (
36
+ AssetSpec(
37
+ name="ESM-C 600M weights",
38
+ filename="esmc_600m_2024_12_v0.pth",
39
+ url=DEFAULT_ESM_URL,
40
+ sha256="8ef856e1a237ee3f995442df997a962e70057faadecf38fc0c8561bd3c2f4324",
41
+ ),
42
+ AssetSpec(
43
+ name="ProtCross checkpoint",
44
+ filename="best-epoch=59.ckpt",
45
+ url=DEFAULT_CHECKPOINT_URL,
46
+ sha256="3eb6d8c9ef94541efc0444508e15d630c156a98e164a6caa08f2ae7a20371e45",
47
+ ),
48
+ AssetSpec(
49
+ name="ProtCross PCA reducer",
50
+ filename="pca_esmc_128.pkl",
51
+ url=DEFAULT_PCA_URL,
52
+ sha256="c4317684fb94c1337a44b844381d7e84472a6958b34a604b0f982984b629098b",
53
+ ),
54
+ )
55
+
56
+
57
+ def get_default_assets_dir() -> Path:
58
+ return Path(os.environ.get("PROTCROSS_ASSETS_DIR", DEFAULT_ASSETS_DIR)).expanduser()
59
+
60
+
61
+ def setup_assets(
62
+ output_dir: str | Path | None = None,
63
+ *,
64
+ esm_url: str = DEFAULT_ESM_URL,
65
+ checkpoint_url: str = DEFAULT_CHECKPOINT_URL,
66
+ pca_url: str = DEFAULT_PCA_URL,
67
+ force: bool = False,
68
+ verify: bool = True,
69
+ skip_esm: bool = False,
70
+ ) -> Path:
71
+ """Download ProtCross assets and return the asset directory."""
72
+ output_dir = Path(output_dir).expanduser() if output_dir else get_default_assets_dir()
73
+ output_dir.mkdir(parents=True, exist_ok=True)
74
+
75
+ specs = [
76
+ AssetSpec(
77
+ DEFAULT_ASSETS[0].name,
78
+ DEFAULT_ASSETS[0].filename,
79
+ esm_url,
80
+ DEFAULT_ASSETS[0].sha256,
81
+ ),
82
+ AssetSpec(
83
+ DEFAULT_ASSETS[1].name,
84
+ DEFAULT_ASSETS[1].filename,
85
+ checkpoint_url,
86
+ DEFAULT_ASSETS[1].sha256,
87
+ ),
88
+ AssetSpec(
89
+ DEFAULT_ASSETS[2].name,
90
+ DEFAULT_ASSETS[2].filename,
91
+ pca_url,
92
+ DEFAULT_ASSETS[2].sha256,
93
+ ),
94
+ ]
95
+ if skip_esm:
96
+ specs = specs[1:]
97
+
98
+ print(f"Installing ProtCross assets into {output_dir}")
99
+ print("Note: ESM-C weights are distributed by EvolutionaryScale under their Hugging Face model terms.")
100
+ for spec in specs:
101
+ download_asset(spec, output_dir / spec.filename, force=force, verify=verify)
102
+
103
+ write_env_file(output_dir, include_esm=(output_dir / DEFAULT_ASSETS[0].filename).exists())
104
+ print("\nAsset setup complete.")
105
+ print("Use with: protcross predict input.pdb --output output.pdb")
106
+ print(f"Environment file written to: {output_dir / 'protcross.env'}")
107
+ return output_dir
108
+
109
+
110
+ def download_asset(spec: AssetSpec, output_path: Path, *, force: bool = False, verify: bool = True) -> None:
111
+ output_path.parent.mkdir(parents=True, exist_ok=True)
112
+ if output_path.exists() and not force:
113
+ if not verify or not spec.sha256 or sha256_file(output_path) == spec.sha256:
114
+ print(f"[skip] {spec.name}: {output_path}")
115
+ return
116
+ print(f"[warn] Existing file failed SHA256 verification and will be replaced: {output_path}")
117
+
118
+ tmp_path = output_path.with_suffix(output_path.suffix + ".part")
119
+ if tmp_path.exists():
120
+ tmp_path.unlink()
121
+
122
+ print(f"[download] {spec.name}")
123
+ print(f" {spec.url}")
124
+ with requests.get(spec.url, stream=True, timeout=30) as response:
125
+ response.raise_for_status()
126
+ total = int(response.headers.get("content-length", 0))
127
+ with tmp_path.open("wb") as file, tqdm(
128
+ total=total if total > 0 else None,
129
+ unit="B",
130
+ unit_scale=True,
131
+ unit_divisor=1024,
132
+ desc=spec.filename,
133
+ ) as progress:
134
+ for chunk in response.iter_content(chunk_size=1024 * 1024):
135
+ if not chunk:
136
+ continue
137
+ file.write(chunk)
138
+ progress.update(len(chunk))
139
+
140
+ if verify and spec.sha256:
141
+ actual = sha256_file(tmp_path)
142
+ if actual != spec.sha256:
143
+ tmp_path.unlink(missing_ok=True)
144
+ raise RuntimeError(
145
+ f"SHA256 mismatch for {spec.filename}: expected {spec.sha256}, got {actual}"
146
+ )
147
+
148
+ tmp_path.replace(output_path)
149
+ print(f"[ok] {output_path}")
150
+
151
+
152
+ def sha256_file(path: str | Path) -> str:
153
+ digest = hashlib.sha256()
154
+ with Path(path).open("rb") as file:
155
+ for chunk in iter(lambda: file.read(1024 * 1024), b""):
156
+ digest.update(chunk)
157
+ return digest.hexdigest()
158
+
159
+
160
+ def write_env_file(output_dir: Path, *, include_esm: bool = True) -> None:
161
+ env_path = output_dir / "protcross.env"
162
+ lines = [
163
+ _export_line("PROTCROSS_ASSETS_DIR", output_dir),
164
+ _export_line("PROTCROSS_CHECKPOINT", output_dir / "best-epoch=59.ckpt"),
165
+ ]
166
+ if include_esm:
167
+ lines.append(_export_line("PROTCROSS_ESM_WEIGHTS", output_dir / "esmc_600m_2024_12_v0.pth"))
168
+ lines.extend(
169
+ [
170
+ _export_line("PROTCROSS_PCA", output_dir / "pca_esmc_128.pkl"),
171
+ "",
172
+ ]
173
+ )
174
+ env_path.write_text("\n".join(lines), encoding="utf-8")
175
+
176
+
177
+ def _export_line(name: str, value: str | Path) -> str:
178
+ return f"export {name}={shlex.quote(str(value))}"
@@ -0,0 +1,2 @@
1
+ """Command line entry points."""
2
+
@@ -0,0 +1,37 @@
1
+ """CLI for downloading matching AlphaFold structures."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ from pathlib import Path
7
+
8
+ from evopoint_da.data.af2 import AF2DownloadConfig, download_af2_structures
9
+
10
+
11
+ def build_parser(prog: str | None = None) -> argparse.ArgumentParser:
12
+ parser = argparse.ArgumentParser(prog=prog, description="Download AlphaFold structures for PDB files.")
13
+ parser.add_argument("--raw-pdb-dir", default="data/raw_pdb")
14
+ parser.add_argument("--output-dir", default="data/raw_af2")
15
+ parser.add_argument("--mapping-file", default="pdb_uniprot_mapping.json")
16
+ parser.add_argument("--max-workers", type=int, default=8)
17
+ parser.add_argument("--uniprot-candidates", type=int, default=3)
18
+ parser.add_argument("--timeout-seconds", type=int, default=30)
19
+ return parser
20
+
21
+
22
+ def main(argv: list[str] | None = None, *, prog: str | None = None) -> int:
23
+ args = build_parser(prog=prog).parse_args(argv)
24
+ config = AF2DownloadConfig(
25
+ raw_pdb_dir=Path(args.raw_pdb_dir),
26
+ output_dir=Path(args.output_dir),
27
+ mapping_file=Path(args.mapping_file),
28
+ max_workers=args.max_workers,
29
+ uniprot_candidates=args.uniprot_candidates,
30
+ timeout_seconds=args.timeout_seconds,
31
+ )
32
+ download_af2_structures(config)
33
+ return 0
34
+
35
+
36
+ if __name__ == "__main__":
37
+ raise SystemExit(main())
@@ -0,0 +1,42 @@
1
+ """Unified ProtCross command line interface."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ from importlib import import_module
7
+
8
+ from evopoint_da import __version__
9
+
10
+
11
+ COMMANDS = {
12
+ "predict": "evopoint_da.cli.predict:main",
13
+ "setup-assets": "evopoint_da.cli.setup_assets:main",
14
+ "preprocess": "evopoint_da.cli.preprocess:main",
15
+ "download-af2": "evopoint_da.cli.download_af2:main",
16
+ "map-labels": "evopoint_da.cli.map_labels:main",
17
+ }
18
+
19
+
20
+ def build_parser() -> argparse.ArgumentParser:
21
+ parser = argparse.ArgumentParser(prog="protcross", description="ProtCross command line tools.")
22
+ parser.add_argument("--version", action="version", version=f"ProtCross {__version__}")
23
+ parser.add_argument("command", choices=sorted(COMMANDS))
24
+ parser.add_argument("args", nargs=argparse.REMAINDER)
25
+ return parser
26
+
27
+
28
+ def main(argv: list[str] | None = None) -> int:
29
+ args = build_parser().parse_args(argv)
30
+ command_args = args.args
31
+ if command_args and command_args[0] == "--":
32
+ command_args = command_args[1:]
33
+ return _load_command(COMMANDS[args.command])(command_args, prog=f"protcross {args.command}")
34
+
35
+
36
+ def _load_command(target: str):
37
+ module_name, function_name = target.split(":", 1)
38
+ return getattr(import_module(module_name), function_name)
39
+
40
+
41
+ if __name__ == "__main__":
42
+ raise SystemExit(main())
@@ -0,0 +1,43 @@
1
+ """CLI for mapping PDB binding labels onto AF2 structures."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ from pathlib import Path
7
+
8
+ from evopoint_da.data.label_mapping import LabelMappingConfig, map_labels
9
+
10
+
11
+ def build_parser(prog: str | None = None) -> argparse.ArgumentParser:
12
+ parser = argparse.ArgumentParser(prog=prog, description="Map PDB-derived labels onto processed AF2 samples.")
13
+ parser.add_argument("--processed-pdb-dir", default="data/processed_pdb")
14
+ parser.add_argument("--processed-af2-dir", default="data/processed_af2")
15
+ parser.add_argument("--raw-pdb-dir", default="data/raw_pdb")
16
+ parser.add_argument("--raw-af2-dir", default="data/raw_af2")
17
+ parser.add_argument("--mapping-file", default="pdb_uniprot_mapping.json")
18
+ parser.add_argument("--output-csv", default="mapping_report_final.csv")
19
+ parser.add_argument("--debug-limit", type=int, default=5)
20
+ parser.add_argument("--min-chain-score", type=float, default=0.15)
21
+ parser.add_argument("--max-rmsd", type=float, default=30.0)
22
+ return parser
23
+
24
+
25
+ def main(argv: list[str] | None = None, *, prog: str | None = None) -> int:
26
+ args = build_parser(prog=prog).parse_args(argv)
27
+ config = LabelMappingConfig(
28
+ processed_pdb_dir=Path(args.processed_pdb_dir),
29
+ processed_af2_dir=Path(args.processed_af2_dir),
30
+ raw_pdb_dir=Path(args.raw_pdb_dir),
31
+ raw_af2_dir=Path(args.raw_af2_dir),
32
+ mapping_file=Path(args.mapping_file),
33
+ output_csv=Path(args.output_csv),
34
+ debug_limit=args.debug_limit,
35
+ min_chain_score=args.min_chain_score,
36
+ max_rmsd=args.max_rmsd,
37
+ )
38
+ map_labels(config)
39
+ return 0
40
+
41
+
42
+ if __name__ == "__main__":
43
+ raise SystemExit(main())
@@ -0,0 +1,156 @@
1
+ """Command line interface for lightweight ProtCross prediction."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ import os
7
+ import sys
8
+ from pathlib import Path
9
+
10
+
11
+ LOCAL_CHECKPOINT = Path("checkpoint/best-epoch=59.ckpt")
12
+ LOCAL_PCA = Path("data/pca_esmc_128.pkl")
13
+
14
+
15
+ def build_parser(prog: str | None = None) -> argparse.ArgumentParser:
16
+ parser = argparse.ArgumentParser(
17
+ prog=prog,
18
+ description="Predict binding-site probabilities for one PDB/mmCIF structure and write them to B-factors.",
19
+ )
20
+ parser.add_argument("input_pdb", nargs="?", help="Input PDB/mmCIF structure.")
21
+ parser.add_argument("--pdb_file", "--pdb-file", dest="pdb_file", help="Legacy input structure argument.")
22
+ parser.add_argument(
23
+ "--assets-dir",
24
+ help=(
25
+ "Directory containing best-epoch=59.ckpt, esmc_600m_2024_12_v0.pth, "
26
+ "and pca_esmc_128.pkl. Explicit file arguments override this."
27
+ ),
28
+ )
29
+ parser.add_argument(
30
+ "--ckpt_path",
31
+ "--ckpt-path",
32
+ "--checkpoint",
33
+ dest="ckpt_path",
34
+ default=os.environ.get("PROTCROSS_CHECKPOINT"),
35
+ help=(
36
+ "ProtCross Lightning checkpoint. Defaults to PROTCROSS_CHECKPOINT, "
37
+ "installed assets, or checkpoint/best-epoch=59.ckpt when present."
38
+ ),
39
+ )
40
+ parser.add_argument(
41
+ "--esm_weights",
42
+ "--esm-weights",
43
+ dest="esm_weights",
44
+ default=os.environ.get("PROTCROSS_ESM_WEIGHTS"),
45
+ help="Local ESM-C 600M weights. Can also be set with PROTCROSS_ESM_WEIGHTS.",
46
+ )
47
+ parser.add_argument(
48
+ "--pca_path",
49
+ "--pca-path",
50
+ "--pca",
51
+ dest="pca_path",
52
+ default=os.environ.get("PROTCROSS_PCA"),
53
+ help=(
54
+ "Fitted PCA pickle for ESM-C embeddings. Defaults to PROTCROSS_PCA, "
55
+ "installed assets, or data/pca_esmc_128.pkl when present."
56
+ ),
57
+ )
58
+ parser.add_argument(
59
+ "-o",
60
+ "--output",
61
+ "--output_pdb",
62
+ "--output-pdb",
63
+ dest="output_pdb",
64
+ help="Output PDB path. Predicted probabilities overwrite B-factors.",
65
+ )
66
+ parser.add_argument("--scores-tsv", help="Optional residue-level score table.")
67
+ parser.add_argument("--threshold", type=float, default=0.5, help="Threshold used in the text summary.")
68
+ parser.add_argument("--chain", "--chain-id", dest="chain_id", help="Restrict prediction to one chain.")
69
+ parser.add_argument("--device", default="auto", help="auto, cpu, cuda or cuda:N.")
70
+ parser.add_argument("--pca_dim", "--pca-dim", dest="pca_dim", type=int, default=128)
71
+ parser.add_argument("--max-len", type=int, default=1022, help="Maximum residues passed to ESM-C.")
72
+ parser.add_argument("--fail-on-truncation", action="store_true")
73
+ parser.add_argument("--quiet", action="store_true")
74
+ return parser
75
+
76
+
77
+ def main(argv: list[str] | None = None, *, prog: str | None = None) -> int:
78
+ parser = build_parser(prog=prog)
79
+ args = parser.parse_args(argv)
80
+
81
+ input_pdb = args.input_pdb or args.pdb_file
82
+ if not input_pdb:
83
+ parser.error("an input PDB/mmCIF path is required")
84
+
85
+ assets = _resolve_asset_directory(args.assets_dir)
86
+ if assets:
87
+ args.ckpt_path = args.ckpt_path or str(assets.checkpoint)
88
+ args.esm_weights = args.esm_weights or str(assets.esm_weights)
89
+ args.pca_path = args.pca_path or str(assets.pca)
90
+
91
+ if not args.ckpt_path and LOCAL_CHECKPOINT.exists():
92
+ args.ckpt_path = str(LOCAL_CHECKPOINT)
93
+ if not args.pca_path and LOCAL_PCA.exists():
94
+ args.pca_path = str(LOCAL_PCA)
95
+
96
+ if not args.esm_weights:
97
+ parser.error(
98
+ "--esm-weights is required unless --assets-dir, PROTCROSS_ESM_WEIGHTS, "
99
+ "or installed assets from `protcross setup-assets` are available"
100
+ )
101
+ if not args.ckpt_path:
102
+ parser.error("--checkpoint is required unless PROTCROSS_CHECKPOINT or installed assets are available")
103
+ if not args.pca_path:
104
+ parser.error("--pca is required unless PROTCROSS_PCA or installed assets are available")
105
+
106
+ try:
107
+ from evopoint_da.inference import ProtCrossPredictor
108
+
109
+ predictor = ProtCrossPredictor.from_files(
110
+ ckpt_path=args.ckpt_path,
111
+ esm_weights=args.esm_weights,
112
+ pca_path=args.pca_path,
113
+ device=args.device,
114
+ pca_dim=args.pca_dim,
115
+ max_len=args.max_len,
116
+ )
117
+ result = predictor.predict(
118
+ input_pdb,
119
+ chain_id=args.chain_id,
120
+ threshold=args.threshold,
121
+ output_pdb=args.output_pdb,
122
+ scores_tsv=args.scores_tsv,
123
+ )
124
+ if args.fail_on_truncation and result.truncated:
125
+ raise RuntimeError(
126
+ f"Input was truncated from {result.original_length} to {len(result.residue_ids)} residues."
127
+ )
128
+ except Exception as exc:
129
+ print(f"ProtCross prediction failed: {exc}", file=sys.stderr)
130
+ return 1
131
+
132
+ if not args.quiet:
133
+ print(result.format_summary())
134
+ if args.output_pdb:
135
+ print(f"Wrote B-factor prediction PDB: {Path(args.output_pdb)}")
136
+ if args.scores_tsv:
137
+ print(f"Wrote score table: {Path(args.scores_tsv)}")
138
+
139
+ return 0
140
+
141
+
142
+ def _resolve_asset_directory(assets_dir: str | None) -> PredictorAssets | None:
143
+ from evopoint_da.inference import PredictorAssets
144
+
145
+ if assets_dir:
146
+ return PredictorAssets.from_dir(assets_dir)
147
+
148
+ default_assets = PredictorAssets.from_default_dir()
149
+ if default_assets.is_complete():
150
+ return default_assets
151
+
152
+ return None
153
+
154
+
155
+ if __name__ == "__main__":
156
+ raise SystemExit(main())
@@ -0,0 +1,45 @@
1
+ """CLI for ESM-C + PCA preprocessing."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ from pathlib import Path
7
+
8
+ from evopoint_da.data.preprocess import PreprocessConfig, preprocess_directory
9
+
10
+
11
+ def build_parser(prog: str | None = None) -> argparse.ArgumentParser:
12
+ parser = argparse.ArgumentParser(prog=prog, description="Preprocess PDB/mmCIF structures into ProtCross .pt files.")
13
+ parser.add_argument("--data_dir", "--data-dir", dest="data_dir", required=True)
14
+ parser.add_argument("--output_dir", "--output-dir", dest="output_dir", required=True)
15
+ parser.add_argument("--model_name", "--model-name", dest="model_name", required=True)
16
+ parser.add_argument("--pca_model_path", "--pca-model-path", dest="pca_model_path", default="pca_esmc_128.pkl")
17
+ parser.add_argument("--fit_pca", "--fit-pca", dest="fit_pca", action="store_true")
18
+ parser.add_argument("--pca_dim", "--pca-dim", dest="pca_dim", type=int, default=128)
19
+ parser.add_argument("--is_af2", "--is-af2", dest="is_af2", action="store_true")
20
+ parser.add_argument("--sample_ratio", "--sample-ratio", dest="sample_ratio", type=float, default=0.1)
21
+ parser.add_argument("--device", default=None)
22
+ parser.add_argument("--max-len", type=int, default=1022)
23
+ return parser
24
+
25
+
26
+ def main(argv: list[str] | None = None, *, prog: str | None = None) -> int:
27
+ args = build_parser(prog=prog).parse_args(argv)
28
+ config = PreprocessConfig(
29
+ data_dir=Path(args.data_dir),
30
+ output_dir=Path(args.output_dir),
31
+ model_name=Path(args.model_name),
32
+ pca_model_path=Path(args.pca_model_path),
33
+ fit_pca=args.fit_pca,
34
+ pca_dim=args.pca_dim,
35
+ is_af2=args.is_af2,
36
+ sample_ratio=args.sample_ratio,
37
+ device=args.device,
38
+ max_len=args.max_len,
39
+ )
40
+ preprocess_directory(config)
41
+ return 0
42
+
43
+
44
+ if __name__ == "__main__":
45
+ raise SystemExit(main())
@@ -0,0 +1,51 @@
1
+ """CLI for downloading ProtCross runtime assets."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+
7
+ from evopoint_da.assets import (
8
+ DEFAULT_CHECKPOINT_URL,
9
+ DEFAULT_ESM_URL,
10
+ DEFAULT_PCA_URL,
11
+ setup_assets,
12
+ )
13
+
14
+
15
+ def build_parser(prog: str | None = None) -> argparse.ArgumentParser:
16
+ parser = argparse.ArgumentParser(
17
+ prog=prog,
18
+ description="Download ProtCross checkpoint, PCA reducer, and ESM-C weights.",
19
+ )
20
+ parser.add_argument(
21
+ "--output-dir",
22
+ help=(
23
+ "Directory where assets will be installed. Defaults to PROTCROSS_ASSETS_DIR "
24
+ "or ~/.cache/protcross/assets/v0.1.1."
25
+ ),
26
+ )
27
+ parser.add_argument("--esm-url", default=DEFAULT_ESM_URL)
28
+ parser.add_argument("--checkpoint-url", default=DEFAULT_CHECKPOINT_URL)
29
+ parser.add_argument("--pca-url", default=DEFAULT_PCA_URL)
30
+ parser.add_argument("--force", action="store_true", help="Re-download files even if they already exist.")
31
+ parser.add_argument("--no-verify", action="store_true", help="Skip SHA256 verification.")
32
+ parser.add_argument("--skip-esm", action="store_true", help="Only download ProtCross checkpoint and PCA reducer.")
33
+ return parser
34
+
35
+
36
+ def main(argv: list[str] | None = None, *, prog: str | None = None) -> int:
37
+ args = build_parser(prog=prog).parse_args(argv)
38
+ setup_assets(
39
+ args.output_dir,
40
+ esm_url=args.esm_url,
41
+ checkpoint_url=args.checkpoint_url,
42
+ pca_url=args.pca_url,
43
+ force=args.force,
44
+ verify=not args.no_verify,
45
+ skip_esm=args.skip_esm,
46
+ )
47
+ return 0
48
+
49
+
50
+ if __name__ == "__main__":
51
+ raise SystemExit(main())
@@ -0,0 +1,19 @@
1
+ """Hydra CLI entry point for standard ProtCross training."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import hydra
6
+ from omegaconf import DictConfig
7
+ from pathlib import Path
8
+
9
+ from evopoint_da.training import run_training
10
+
11
+ CONFIG_PATH = str(Path(__file__).resolve().parents[3] / "configs")
12
+
13
+ @hydra.main(version_base="1.3", config_path=CONFIG_PATH, config_name="train")
14
+ def main(cfg: DictConfig) -> None:
15
+ run_training(cfg)
16
+
17
+
18
+ if __name__ == "__main__":
19
+ main()
@@ -0,0 +1,19 @@
1
+ """Data parsing and loading utilities."""
2
+
3
+ from .af2 import AF2DownloadConfig, AF2Downloader, download_af2_structures
4
+ from .label_mapping import LabelMappingConfig, map_labels
5
+ from .pca import PCAReducer
6
+ from .structure import MAX_ESM_RESIDUES, STANDARD_AA, StructureParser, truncate_parsed_structure
7
+
8
+ __all__ = [
9
+ "AF2DownloadConfig",
10
+ "AF2Downloader",
11
+ "LabelMappingConfig",
12
+ "MAX_ESM_RESIDUES",
13
+ "PCAReducer",
14
+ "STANDARD_AA",
15
+ "StructureParser",
16
+ "download_af2_structures",
17
+ "map_labels",
18
+ "truncate_parsed_structure",
19
+ ]