fullFold 0.1.0__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.
fullFold/__init__.py ADDED
@@ -0,0 +1,3 @@
1
+ """Wrapper-only multi-GPU scheduling layer around AlphaFold 3."""
2
+
3
+ __version__ = '0.1.0'
fullFold/__main__.py ADDED
@@ -0,0 +1,4 @@
1
+ from fullFold.cli import main
2
+
3
+ if __name__ == '__main__':
4
+ raise SystemExit(main())
fullFold/banner.py ADDED
@@ -0,0 +1,20 @@
1
+ """Startup banner for the fullFold CLI. Printed to stderr."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import sys
6
+
7
+ BANNER = """\
8
+ __ _ _ _____ _ _
9
+ / _|_ _| | | ___|__ | | __| |
10
+ | |_| | | | | | |_ / _ \\| |/ _` |
11
+ | _| |_| | | | _| (_) | | (_| |
12
+ |_| \\__,_|_|_|_| \\___/|_|\\__,_|
13
+
14
+ fullFold: unlocking the full speed of AF3!
15
+ Please cite AlphaFold3 and fullFold (see our repository).
16
+ """
17
+
18
+
19
+ def print_banner(file=None) -> None:
20
+ print(BANNER, file=file or sys.stderr, end='')
fullFold/benchmark.py ADDED
@@ -0,0 +1,232 @@
1
+ """GPU discovery and 1024-token compile/throughput probe."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import os
7
+ import subprocess
8
+ import sys
9
+ import warnings
10
+ from dataclasses import asdict, dataclass
11
+ from pathlib import Path
12
+
13
+ from fullFold.config import Config, host_id, worker_environ, write_atomic
14
+ from fullFold.scheduling import compilation_modifier, inference_modifier
15
+
16
+ T_REF_S = 59.423 # A100 tokamax inference-only seconds at bucket 1024
17
+
18
+
19
+ @dataclass(frozen=True)
20
+ class Gpu:
21
+ slot: int
22
+ physical_id: str
23
+ uuid: str
24
+ pci_bus_id: str
25
+ device_kind: str
26
+ memory_bytes: int
27
+
28
+
29
+ @dataclass(frozen=True)
30
+ class Bench:
31
+ t_ms: tuple[float, float, float, float]
32
+ s_ms: float
33
+ r_ms: float
34
+ contaminated: bool
35
+ host: str
36
+ gpu_key: str
37
+
38
+
39
+ def gpu_multiplier(bench: Bench) -> float:
40
+ """m_gpu = S_1024 / T_REF from the probe's inference-only time."""
41
+ s_s = bench.s_ms / 1000.0
42
+ if T_REF_S <= 0 or s_s <= 0:
43
+ return 1.0
44
+ return s_s / T_REF_S
45
+
46
+
47
+ def predict_inference_s(m_gpu: float, bucket: int) -> float:
48
+ """T_predicted = m_gpu * inference_modifier(bucket) * T_REF, in seconds."""
49
+ return m_gpu * inference_modifier(bucket) * T_REF_S
50
+
51
+
52
+ def compile_overhead_ms(bench: Bench, shape: int) -> float:
53
+ """Probe R at 1024 scaled by the CSV compile modifier (last-row above 5216)."""
54
+ return bench.r_ms * compilation_modifier(shape)
55
+
56
+
57
+ def summarise_timings(t_ms: tuple[float, ...] | list[float]) -> tuple[float, float, bool]:
58
+ """Tokamax pays compile on seeds 1 and 2; seeds 3–4 are steady-state inference.
59
+
60
+ ``contaminated`` means S is unusable (t3 vs t4 disagree). Negative R is a
61
+ warm cache or inverted compile, not a reason to ignore throughput.
62
+ """
63
+ t = tuple(float(x) for x in t_ms)
64
+ s = (t[2] + t[3]) / 2.0
65
+ r = t[0] + t[1] - 2.0 * s
66
+ contaminated = s > 0 and abs(t[2] - t[3]) / s > 0.25
67
+ return s, r, contaminated
68
+
69
+
70
+ def nvidia_smi() -> list[dict]:
71
+ cmd = [
72
+ 'nvidia-smi',
73
+ '--query-gpu=index,uuid,pci.bus_id,name,memory.total',
74
+ '--format=csv,noheader,nounits',
75
+ ]
76
+ try:
77
+ out = subprocess.check_output(cmd, text=True, stderr=subprocess.DEVNULL)
78
+ except (OSError, subprocess.CalledProcessError):
79
+ return []
80
+ rows = []
81
+ for line in out.splitlines():
82
+ parts = [p.strip() for p in line.split(',')]
83
+ if len(parts) < 5:
84
+ continue
85
+ try:
86
+ mem = int(float(parts[4])) * 1024 * 1024
87
+ except ValueError:
88
+ mem = 0
89
+ rows.append({
90
+ 'index': parts[0], 'uuid': parts[1], 'pci': parts[2],
91
+ 'name': parts[3], 'memory_bytes': mem,
92
+ })
93
+ rows.sort(key=lambda r: r['pci'])
94
+ return rows
95
+
96
+
97
+ def discover_gpus(cfg: Config, rows: list[dict] | None = None) -> list[Gpu]:
98
+ rows = list(rows if rows is not None else nvidia_smi())
99
+ env = os.environ.get('CUDA_VISIBLE_DEVICES')
100
+ wanted = list(cfg.gpus) if cfg.gpus else (
101
+ [x.strip() for x in env.split(',') if x.strip()]
102
+ if env not in (None, '') else []
103
+ )
104
+ if wanted:
105
+ selected = []
106
+ for w in wanted:
107
+ hit = next((r for r in rows if r['index'] == w or r['uuid'] == w), None)
108
+ if hit is None and rows:
109
+ raise ValueError(f'GPU {w!r} is not in nvidia-smi output')
110
+ if hit:
111
+ selected.append(hit)
112
+ rows = selected
113
+ gpus = [
114
+ Gpu(slot=i, physical_id=r['index'], uuid=r['uuid'],
115
+ pci_bus_id=r['pci'], device_kind=r['name'],
116
+ memory_bytes=int(r['memory_bytes']))
117
+ for i, r in enumerate(rows)
118
+ ]
119
+ if not gpus:
120
+ raise RuntimeError('no GPUs discovered (nvidia-smi empty and no --gpus)')
121
+ return gpus
122
+
123
+
124
+ def gpu_key(gpu: Gpu, jax_version: str = '', af3_version: str = '') -> str:
125
+ return f'{gpu.pci_bus_id}|{gpu.device_kind}|{jax_version}|{af3_version}'
126
+
127
+
128
+ def cache_path(cfg: Config, gpu: Gpu, jax_version: str = '', af3_version: str = '') -> Path:
129
+ return cfg.cache_dir / f'{host_id()}__{gpu_key(gpu, jax_version, af3_version)}.json'
130
+
131
+
132
+ def load_bench(path: Path) -> Bench | None:
133
+ if not path.is_file():
134
+ return None
135
+ d = json.loads(path.read_text())
136
+ t = tuple(d['t_ms'])
137
+ s, r, cont = summarise_timings(t)
138
+ return Bench(t_ms=t, s_ms=s, r_ms=r, contaminated=cont,
139
+ host=d.get('host', ''), gpu_key=d.get('gpu_key', ''))
140
+
141
+
142
+ def save_bench(path: Path, bench: Bench) -> None:
143
+ write_atomic(path, json.dumps(asdict(bench), indent=2))
144
+
145
+
146
+ def build_probe_json(n_tokens: int = 1024, rng_seed: int = 42) -> dict:
147
+ import random
148
+ aa = list('ACDEFGHIKLMNPQRSTVWY')
149
+ rng = random.Random(rng_seed)
150
+ seq = ''.join(rng.choices(aa, k=n_tokens))
151
+ msa = [f'>query\n{seq}\n']
152
+ for i in range(31):
153
+ msa.append(f'>seq_{i+1}\n' + ''.join(rng.choices(aa, k=n_tokens)) + '\n')
154
+ return {
155
+ 'name': f'fullFold_bench_{n_tokens}',
156
+ 'modelSeeds': [0, 1, 2, 3],
157
+ 'sequences': [{
158
+ 'protein': {
159
+ 'id': 'A', 'sequence': seq,
160
+ 'unpairedMsa': ''.join(msa), 'pairedMsa': '', 'templates': [],
161
+ }
162
+ }],
163
+ 'dialect': 'alphafold3', 'version': 3,
164
+ }
165
+
166
+
167
+ def to_scheduler_gpu(bench: Bench) -> tuple[int, int]:
168
+ """(compile_us, infer_us_1024) for scheduling.py.
169
+
170
+ ``m_gpu = S_1024 / T_REF`` from the probe; inference at bucket b is
171
+ ``m_gpu * f(b/1024) * T_REF``.
172
+ """
173
+ compile_us = max(int(round(compile_overhead_ms(bench, 1024) * 1000)), 0)
174
+ infer_us_1024 = max(1, int(round(gpu_multiplier(bench) * T_REF_S * 1e6)))
175
+ return compile_us, infer_us_1024
176
+
177
+
178
+ def measure(gpu: Gpu, cfg: Config) -> Bench:
179
+ """Spawn a per-GPU probe subprocess so the parent never initialises JAX."""
180
+ import sys
181
+ env = worker_environ(
182
+ cfg, gpu.physical_id, pci_bus_id=gpu.pci_bus_id, device_kind=gpu.device_kind)
183
+ cmd = [
184
+ sys.executable, '-m', 'fullFold.worker', '--probe',
185
+ '--gpu', str(gpu.physical_id),
186
+ '--model-dir', str(cfg.model_dir),
187
+ '--output-dir', str(cfg.output_dir),
188
+ '--flash-attention', cfg.flash_attention,
189
+ ]
190
+ r = subprocess.run(cmd, env=env, capture_output=True, text=True)
191
+ if r.returncode != 0:
192
+ detail = (r.stderr or r.stdout or '(no output)').strip()
193
+ raise RuntimeError(
194
+ f'probe on GPU {gpu.physical_id} exited {r.returncode}:\n{detail}'
195
+ )
196
+ t = json.loads(r.stdout.strip().splitlines()[-1])
197
+ s, r_ms, cont = summarise_timings(t)
198
+ if cont:
199
+ warnings.warn(f'contaminated benchmark on GPU {gpu.physical_id}: t={t}')
200
+ b = Bench(t_ms=tuple(t), s_ms=s, r_ms=r_ms, contaminated=cont,
201
+ host=host_id(), gpu_key=gpu_key(gpu))
202
+ save_bench(cache_path(cfg, gpu), b)
203
+ return b
204
+
205
+
206
+ def get_or_measure(
207
+ gpus: list[Gpu], cfg: Config, measure_fn=None,
208
+ ) -> list[tuple[Gpu, Bench]]:
209
+ measure_fn = measure_fn or measure
210
+ out: list[tuple[Gpu, Bench]] = []
211
+ missing: list[Gpu] = []
212
+ for g in gpus:
213
+ p = cache_path(cfg, g)
214
+ b = None if cfg.force_benchmark else load_bench(p)
215
+ if b is None:
216
+ missing.append(g)
217
+ else:
218
+ out.append((g, b))
219
+ if missing:
220
+ n = len(missing)
221
+ gpu_word = 'GPU' if n == 1 else 'GPUs'
222
+ print(
223
+ f'A few-minute benchmark of the available GPUs is running '
224
+ f'({n} {gpu_word}).',
225
+ file=sys.stderr,
226
+ )
227
+ from concurrent.futures import ThreadPoolExecutor
228
+ with ThreadPoolExecutor(max_workers=n) as ex:
229
+ benches = list(ex.map(lambda g: measure_fn(g, cfg), missing))
230
+ out.extend(zip(missing, benches))
231
+ out.sort(key=lambda x: x[0].slot)
232
+ return out
fullFold/cli.py ADDED
@@ -0,0 +1,103 @@
1
+ """argparse -> Config -> one function. No scheduling logic here."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import argparse
6
+ from pathlib import Path
7
+
8
+ from fullFold.banner import print_banner
9
+ from fullFold.config import DEFAULT_BUCKETS, Config, config_from_dict
10
+
11
+
12
+ def _buckets(s: str) -> tuple[int, ...]:
13
+ return tuple(int(x) for x in s.split(',')) if s else DEFAULT_BUCKETS
14
+
15
+
16
+ def _add_io(p: argparse.ArgumentParser) -> None:
17
+ p.add_argument('--input-dir', type=Path, required=True)
18
+ p.add_argument('--output-dir', type=Path, required=True)
19
+
20
+
21
+ def _add_sched(p: argparse.ArgumentParser) -> None:
22
+ p.add_argument('--model-dir', type=Path, default=Config.model_dir)
23
+ p.add_argument('--gpus', default='')
24
+ p.add_argument('--no-prefetch', dest='prefetch', action='store_const',
25
+ const=0, default=1)
26
+ p.add_argument('--no-background-extract', dest='background_extract',
27
+ action='store_false')
28
+ p.add_argument('--policy', choices=('contiguous', 'roundrobin'), default='contiguous')
29
+ p.add_argument('--bucket-mode', choices=('ladder', 'free'), default='free')
30
+ p.add_argument('--buckets', type=_buckets, default=DEFAULT_BUCKETS)
31
+ p.add_argument('--bucket-margin', type=float, default=0.05)
32
+ p.add_argument('--stale-lock-seconds', type=int, default=900)
33
+ p.add_argument('--retry-failed', action='store_true')
34
+ p.add_argument('--exact-split-threshold', type=int, default=512)
35
+ p.add_argument('--force-benchmark', action='store_true')
36
+ p.add_argument('--bench-seed', type=int, default=42)
37
+ p.add_argument('--cache-dir', type=Path, default=Config.cache_dir)
38
+ p.add_argument('--exact-tokens', action='store_true')
39
+ p.add_argument('--dry-run', action='store_true')
40
+ p.add_argument('--jax-compilation-cache-dir', type=Path, default=None)
41
+ p.add_argument('--xla-mem-fraction', type=float, default=0.97)
42
+ p.add_argument('--save-embeddings', action='store_true')
43
+ p.add_argument('--save-distogram', action='store_true')
44
+ p.add_argument('--compress-large-output-files', action='store_true')
45
+ p.add_argument('--num-recycles', type=int, default=10)
46
+ p.add_argument('--num-diffusion-samples', type=int, default=5)
47
+ p.add_argument('--flash-attention', default='triton')
48
+
49
+
50
+ def _cfg(ns: argparse.Namespace) -> Config:
51
+ kw = {n: getattr(ns, n) for n in Config.__dataclass_fields__ if hasattr(ns, n)}
52
+ g = kw.get('gpus', ())
53
+ if isinstance(g, str):
54
+ kw['gpus'] = tuple(x for x in g.split(',') if x)
55
+ return config_from_dict(kw)
56
+
57
+
58
+ def main(argv: list[str] | None = None) -> int:
59
+ print_banner()
60
+ parser = argparse.ArgumentParser(prog='fullFold')
61
+ sub = parser.add_subparsers(dest='cmd', required=True)
62
+
63
+ for name in ('scan', 'plan', 'run'):
64
+ p = sub.add_parser(name)
65
+ _add_io(p)
66
+ _add_sched(p)
67
+ if name == 'run':
68
+ p.add_argument('-v', '--verbose', action='store_true',
69
+ help='Show per-GPU progress bars during the run')
70
+
71
+ p = sub.add_parser('benchmark')
72
+ p.add_argument('--output-dir', type=Path, default=Path('.'))
73
+ p.add_argument('--input-dir', type=Path, default=Path('.'))
74
+ _add_sched(p)
75
+
76
+ p = sub.add_parser('template')
77
+ p.add_argument('--template', type=Path, required=True)
78
+ p.add_argument('--records', type=Path, required=True)
79
+ p.add_argument('--output-dir', type=Path, required=True)
80
+ p.add_argument('--type', dest='record_type', default='protein')
81
+
82
+ args = parser.parse_args(argv)
83
+ cfg = _cfg(args)
84
+ if args.cmd == 'scan':
85
+ from fullFold.engine import cmd_scan
86
+ return cmd_scan(cfg)
87
+ if args.cmd == 'benchmark':
88
+ from fullFold.engine import cmd_benchmark
89
+ return cmd_benchmark(cfg)
90
+ if args.cmd == 'plan':
91
+ from fullFold.engine import cmd_plan
92
+ return cmd_plan(cfg)
93
+ if args.cmd == 'run':
94
+ from fullFold.engine import cmd_run
95
+ return cmd_run(cfg)
96
+ if args.cmd == 'template':
97
+ from fullFold.templates import cmd_template
98
+ return cmd_template(cfg)
99
+ return 2
100
+
101
+
102
+ if __name__ == '__main__':
103
+ raise SystemExit(main())
fullFold/config.py ADDED
@@ -0,0 +1,148 @@
1
+ """Single frozen Config plus hashing helpers."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import dataclasses
6
+ import hashlib
7
+ import json
8
+ import os
9
+ import socket
10
+ from pathlib import Path
11
+
12
+ DEFAULT_BUCKETS = (
13
+ 128, 256, 384, 512, 768, 1024, 1280, 1536, 2048, 2560, 3072, 3584, 4096,
14
+ 4608, 5120,
15
+ )
16
+
17
+
18
+ @dataclasses.dataclass(frozen=True)
19
+ class Config:
20
+ input_dir: Path = Path('.')
21
+ output_dir: Path = Path('.')
22
+ model_dir: Path = Path('~/models').expanduser()
23
+ gpus: tuple[str, ...] = ()
24
+ dry_run: bool = False
25
+ verbose: bool = False
26
+ prefetch: int = 1
27
+ background_extract: bool = True
28
+ policy: str = 'contiguous' # contiguous | roundrobin
29
+ bucket_mode: str = 'free' # free | ladder
30
+ buckets: tuple[int, ...] = DEFAULT_BUCKETS
31
+ bucket_margin: float = 0.05
32
+ stale_lock_seconds: int = 900
33
+ retry_failed: bool = False
34
+ exact_split_threshold: int = 512
35
+ force_benchmark: bool = False
36
+ bench_seed: int = 42
37
+ cache_dir: Path = Path('~/.cache/fullFold/bench').expanduser()
38
+ jax_compilation_cache_dir: Path | None = None
39
+ xla_mem_fraction: float = 0.97
40
+ xla_preallocate: bool = True
41
+ exact_tokens: bool = False
42
+ template: Path | None = None
43
+ records: Path | None = None
44
+ record_type: str = 'protein'
45
+ save_embeddings: bool = False
46
+ save_distogram: bool = False
47
+ compress_large_output_files: bool = False
48
+ save_terms_of_use: bool = True
49
+ num_recycles: int = 10
50
+ num_diffusion_samples: int = 5
51
+ flash_attention: str = 'triton'
52
+
53
+
54
+ def host_id() -> str:
55
+ p = Path('/etc/machine-id')
56
+ if p.is_file():
57
+ return p.read_text().strip()
58
+ return hashlib.sha256(socket.gethostname().encode()).hexdigest()[:16]
59
+
60
+
61
+ def sha256_bytes(data: bytes) -> str:
62
+ return hashlib.sha256(data).hexdigest()
63
+
64
+
65
+ def sha256_file(path: Path) -> str:
66
+ h = hashlib.sha256()
67
+ with path.open('rb') as f:
68
+ for chunk in iter(lambda: f.read(1 << 16), b''):
69
+ h.update(chunk)
70
+ return h.hexdigest()
71
+
72
+
73
+ def sanitise(name: str) -> str:
74
+ allowed = set('abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789_-.')
75
+ return ''.join(c for c in name.replace(' ', '_') if c in allowed)
76
+
77
+
78
+ def config_hash(cfg: Config) -> str:
79
+ return sha256_bytes(json.dumps(config_to_dict(cfg), sort_keys=True).encode())[:16]
80
+
81
+
82
+ def config_to_dict(cfg: Config) -> dict:
83
+ d = dataclasses.asdict(cfg)
84
+ for k, v in d.items():
85
+ if isinstance(v, Path):
86
+ d[k] = str(v)
87
+ elif isinstance(v, tuple):
88
+ d[k] = list(v)
89
+ return d
90
+
91
+
92
+ def config_from_dict(d: dict) -> Config:
93
+ fields = Config.__dataclass_fields__
94
+ kw = {}
95
+ for k, v in d.items():
96
+ if k not in fields:
97
+ continue
98
+ hint = str(fields[k].type)
99
+ if v is None:
100
+ kw[k] = None
101
+ elif 'Path' in hint and not isinstance(v, Path):
102
+ kw[k] = Path(v)
103
+ elif 'tuple' in hint and not isinstance(v, tuple):
104
+ kw[k] = tuple(v)
105
+ else:
106
+ kw[k] = v
107
+ return Config(**kw)
108
+
109
+
110
+ def write_atomic(path: Path, text: str) -> None:
111
+ path.parent.mkdir(parents=True, exist_ok=True)
112
+ tmp = path.with_name(path.name + '.tmp')
113
+ tmp.write_text(text)
114
+ os.replace(tmp, path)
115
+
116
+
117
+ def jax_cache_dir(
118
+ cfg: Config, *, physical_id: str, pci_bus_id: str = '', device_kind: str = '',
119
+ ) -> Path:
120
+ """Persistent XLA cache root for this host + GPU name. Always enabled.
121
+
122
+ PCI / physical_id are accepted for call-site compatibility but are not
123
+ part of the path: identical nvidia-smi names share a directory. JAX's
124
+ file keys (HLO, jaxlib, GPU-name topology) still miss across models.
125
+ """
126
+ del physical_id, pci_bus_id
127
+ root = Path(cfg.jax_compilation_cache_dir) if cfg.jax_compilation_cache_dir else (
128
+ cfg.cache_dir / 'jax')
129
+ kind_part = sanitise(device_kind) if device_kind else 'gpu'
130
+ return root / f'{host_id()}__{kind_part}'
131
+
132
+
133
+ def worker_environ(
134
+ cfg: Config, gpu_physical_id: str, *,
135
+ pci_bus_id: str = '', device_kind: str = '',
136
+ ) -> dict[str, str]:
137
+ """Env for a one-GPU worker/probe subprocess (parent never initialises JAX)."""
138
+ env = os.environ.copy()
139
+ env['CUDA_VISIBLE_DEVICES'] = str(gpu_physical_id)
140
+ env['CUDA_DEVICE_ORDER'] = 'PCI_BUS_ID'
141
+ env['XLA_CLIENT_MEM_FRACTION'] = str(cfg.xla_mem_fraction)
142
+ env['XLA_PYTHON_CLIENT_PREALLOCATE'] = 'true' if cfg.xla_preallocate else 'false'
143
+ cache = jax_cache_dir(
144
+ cfg, physical_id=gpu_physical_id, pci_bus_id=pci_bus_id,
145
+ device_kind=device_kind)
146
+ cache.mkdir(parents=True, exist_ok=True)
147
+ env['AF3SCHED_JAX_CACHE'] = str(cache)
148
+ return env