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 +3 -0
- fullFold/__main__.py +4 -0
- fullFold/banner.py +20 -0
- fullFold/benchmark.py +232 -0
- fullFold/cli.py +103 -0
- fullFold/config.py +148 -0
- fullFold/data/reference_timings.csv +312 -0
- fullFold/engine.py +476 -0
- fullFold/jobs.py +186 -0
- fullFold/runner.py +118 -0
- fullFold/scheduling.py +506 -0
- fullFold/templates.py +233 -0
- fullFold/tokens.py +155 -0
- fullFold/worker.py +615 -0
- fullfold-0.1.0.dist-info/METADATA +193 -0
- fullfold-0.1.0.dist-info/RECORD +20 -0
- fullfold-0.1.0.dist-info/WHEEL +5 -0
- fullfold-0.1.0.dist-info/entry_points.txt +3 -0
- fullfold-0.1.0.dist-info/licenses/LICENSE +201 -0
- fullfold-0.1.0.dist-info/top_level.txt +1 -0
fullFold/__init__.py
ADDED
fullFold/__main__.py
ADDED
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
|