boltz2-python-client 0.2__py3-none-any.whl → 0.3.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.
- boltz2_client/__init__.py +28 -1
- boltz2_client/__main__.py +9 -0
- boltz2_client/cli.py +456 -44
- boltz2_client/client.py +539 -16
- boltz2_client/models.py +1 -1
- boltz2_client/msa_search.py +489 -0
- boltz2_client/multi_endpoint_client.py +1679 -0
- boltz2_client/virtual_screening.py +11 -3
- {boltz2_python_client-0.2.dist-info → boltz2_python_client-0.3.0.dist-info}/METADATA +123 -31
- boltz2_python_client-0.3.0.dist-info/RECORD +17 -0
- boltz2_python_client-0.2.dist-info/RECORD +0 -14
- {boltz2_python_client-0.2.dist-info → boltz2_python_client-0.3.0.dist-info}/WHEEL +0 -0
- {boltz2_python_client-0.2.dist-info → boltz2_python_client-0.3.0.dist-info}/entry_points.txt +0 -0
- {boltz2_python_client-0.2.dist-info → boltz2_python_client-0.3.0.dist-info}/licenses/LICENSE +0 -0
- {boltz2_python_client-0.2.dist-info → boltz2_python_client-0.3.0.dist-info}/top_level.txt +0 -0
boltz2_client/__init__.py
CHANGED
|
@@ -26,8 +26,9 @@ Example:
|
|
|
26
26
|
>>> print(f"Confidence: {result.confidence_scores[0]:.3f}")
|
|
27
27
|
"""
|
|
28
28
|
|
|
29
|
-
__version__ = "0.
|
|
29
|
+
__version__ = "0.3.0"
|
|
30
30
|
__author__ = "NVIDIA Corporation"
|
|
31
|
+
__email__ = "bionemo-support@nvidia.com"
|
|
31
32
|
|
|
32
33
|
from .client import Boltz2Client, Boltz2SyncClient, EndpointType
|
|
33
34
|
from .models import (
|
|
@@ -39,6 +40,7 @@ from .models import (
|
|
|
39
40
|
BondConstraint,
|
|
40
41
|
Atom,
|
|
41
42
|
AlignmentFileRecord,
|
|
43
|
+
AlignmentFormat,
|
|
42
44
|
HealthStatus,
|
|
43
45
|
ServiceMetadata,
|
|
44
46
|
)
|
|
@@ -57,6 +59,18 @@ from .virtual_screening import (
|
|
|
57
59
|
VirtualScreeningResult,
|
|
58
60
|
quick_screen,
|
|
59
61
|
)
|
|
62
|
+
from .multi_endpoint_client import (
|
|
63
|
+
MultiEndpointClient,
|
|
64
|
+
LoadBalanceStrategy,
|
|
65
|
+
EndpointConfig,
|
|
66
|
+
)
|
|
67
|
+
from .msa_search import (
|
|
68
|
+
MSASearchClient,
|
|
69
|
+
MSASearchIntegration,
|
|
70
|
+
MSASearchRequest,
|
|
71
|
+
MSASearchResponse,
|
|
72
|
+
MSAFormatConverter,
|
|
73
|
+
)
|
|
60
74
|
|
|
61
75
|
# Optional imports for visualization
|
|
62
76
|
try:
|
|
@@ -95,6 +109,7 @@ __all__ = [
|
|
|
95
109
|
"BondConstraint",
|
|
96
110
|
"Atom",
|
|
97
111
|
"AlignmentFileRecord",
|
|
112
|
+
"AlignmentFormat",
|
|
98
113
|
"HealthStatus",
|
|
99
114
|
"ServiceMetadata",
|
|
100
115
|
"AffinityPrediction",
|
|
@@ -112,6 +127,18 @@ __all__ = [
|
|
|
112
127
|
"CompoundLibrary",
|
|
113
128
|
"VirtualScreeningResult",
|
|
114
129
|
"quick_screen",
|
|
130
|
+
|
|
131
|
+
# Multi-endpoint support
|
|
132
|
+
"MultiEndpointClient",
|
|
133
|
+
"LoadBalanceStrategy",
|
|
134
|
+
"EndpointConfig",
|
|
135
|
+
|
|
136
|
+
# MSA Search NIM integration
|
|
137
|
+
"MSASearchClient",
|
|
138
|
+
"MSASearchIntegration",
|
|
139
|
+
"MSASearchRequest",
|
|
140
|
+
"MSASearchResponse",
|
|
141
|
+
"MSAFormatConverter",
|
|
115
142
|
]
|
|
116
143
|
|
|
117
144
|
# Add visualization exports if available
|
boltz2_client/cli.py
CHANGED
|
@@ -21,7 +21,7 @@ from typing import List, Optional, Tuple, Dict, Any
|
|
|
21
21
|
import click
|
|
22
22
|
from rich.console import Console
|
|
23
23
|
from rich.table import Table
|
|
24
|
-
from rich.progress import Progress, SpinnerColumn, TextColumn, TimeElapsedColumn
|
|
24
|
+
from rich.progress import Progress, SpinnerColumn, TextColumn, TimeElapsedColumn, BarColumn
|
|
25
25
|
from rich.panel import Panel
|
|
26
26
|
from rich.text import Text
|
|
27
27
|
import yaml as pyyaml
|
|
@@ -58,7 +58,7 @@ def print_warning(message: str):
|
|
|
58
58
|
|
|
59
59
|
|
|
60
60
|
@click.group()
|
|
61
|
-
@click.option('--base-url', default='http://localhost:8000', help='Service base URL')
|
|
61
|
+
@click.option('--base-url', default='http://localhost:8000', help='Service base URL (can be comma-separated for multiple endpoints)')
|
|
62
62
|
@click.option('--api-key', help='API key for NVIDIA hosted endpoints (or set NVIDIA_API_KEY env var)')
|
|
63
63
|
@click.option('--endpoint-type',
|
|
64
64
|
type=click.Choice(['local', 'nvidia_hosted']),
|
|
@@ -66,10 +66,16 @@ def print_warning(message: str):
|
|
|
66
66
|
help='Type of endpoint: local or nvidia_hosted')
|
|
67
67
|
@click.option('--timeout', default=300.0, help='Request timeout in seconds')
|
|
68
68
|
@click.option('--poll-seconds', default=10, help='Polling interval for NVIDIA hosted endpoints')
|
|
69
|
+
@click.option('--multi-endpoint', is_flag=True, help='Enable multi-endpoint load balancing')
|
|
70
|
+
@click.option('--load-balance-strategy',
|
|
71
|
+
type=click.Choice(['round_robin', 'least_loaded', 'random']),
|
|
72
|
+
default='least_loaded',
|
|
73
|
+
help='Load balancing strategy for multi-endpoint')
|
|
69
74
|
@click.option('--verbose', '-v', is_flag=True, help='Enable verbose output')
|
|
70
75
|
@click.pass_context
|
|
71
76
|
def cli(ctx, base_url: str, api_key: Optional[str], endpoint_type: str,
|
|
72
|
-
timeout: float, poll_seconds: int,
|
|
77
|
+
timeout: float, poll_seconds: int, multi_endpoint: bool,
|
|
78
|
+
load_balance_strategy: str, verbose: bool):
|
|
73
79
|
"""
|
|
74
80
|
Boltz-2 Python Client CLI
|
|
75
81
|
|
|
@@ -80,6 +86,12 @@ def cli(ctx, base_url: str, api_key: Optional[str], endpoint_type: str,
|
|
|
80
86
|
# Local endpoint
|
|
81
87
|
boltz2 --base-url http://localhost:8000 protein "MKTVRQERLK..."
|
|
82
88
|
|
|
89
|
+
# Multi-endpoint for parallel processing
|
|
90
|
+
boltz2 --multi-endpoint --base-url "http://localhost:8000,http://localhost:8001,http://localhost:8002,http://localhost:8003" screen target.fasta compounds.csv
|
|
91
|
+
|
|
92
|
+
# Multi-endpoint with custom strategy
|
|
93
|
+
boltz2 --multi-endpoint --base-url "http://localhost:8000,http://localhost:8001" --load-balance-strategy round_robin protein "MKTVRQERLK..."
|
|
94
|
+
|
|
83
95
|
# NVIDIA hosted endpoint
|
|
84
96
|
boltz2 --base-url https://health.api.nvidia.com --endpoint-type nvidia_hosted --api-key YOUR_KEY protein "MKTVRQERLK..."
|
|
85
97
|
|
|
@@ -93,27 +105,61 @@ def cli(ctx, base_url: str, api_key: Optional[str], endpoint_type: str,
|
|
|
93
105
|
ctx.obj['endpoint_type'] = endpoint_type
|
|
94
106
|
ctx.obj['timeout'] = timeout
|
|
95
107
|
ctx.obj['poll_seconds'] = poll_seconds
|
|
108
|
+
ctx.obj['multi_endpoint'] = multi_endpoint
|
|
109
|
+
ctx.obj['load_balance_strategy'] = load_balance_strategy
|
|
96
110
|
ctx.obj['verbose'] = verbose
|
|
97
111
|
|
|
98
112
|
if verbose:
|
|
99
|
-
|
|
100
|
-
|
|
101
|
-
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
print_info("
|
|
113
|
+
if multi_endpoint:
|
|
114
|
+
endpoints = [url.strip() for url in base_url.split(',')]
|
|
115
|
+
print_info(f"Using multi-endpoint mode with {len(endpoints)} endpoints")
|
|
116
|
+
print_info(f"Load balance strategy: {load_balance_strategy}")
|
|
117
|
+
for ep in endpoints:
|
|
118
|
+
print_info(f" - {ep}")
|
|
119
|
+
else:
|
|
120
|
+
print_info(f"Using {endpoint_type} endpoint: {base_url}")
|
|
121
|
+
if endpoint_type == 'nvidia_hosted':
|
|
122
|
+
if api_key:
|
|
123
|
+
print_info("API key provided via command line")
|
|
124
|
+
else:
|
|
125
|
+
print_info("API key will be read from NVIDIA_API_KEY environment variable")
|
|
105
126
|
|
|
106
127
|
|
|
107
|
-
def create_client(ctx)
|
|
108
|
-
"""Create a Boltz2Client from context."""
|
|
109
|
-
|
|
110
|
-
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
|
|
115
|
-
|
|
116
|
-
|
|
128
|
+
def create_client(ctx):
|
|
129
|
+
"""Create a Boltz2Client or MultiEndpointClient from context."""
|
|
130
|
+
from .multi_endpoint_client import MultiEndpointClient, LoadBalanceStrategy
|
|
131
|
+
|
|
132
|
+
if ctx.obj['multi_endpoint']:
|
|
133
|
+
# Parse multiple endpoints from comma-separated list
|
|
134
|
+
endpoints = [url.strip() for url in ctx.obj['base_url'].split(',')]
|
|
135
|
+
|
|
136
|
+
# Map strategy string to enum
|
|
137
|
+
strategy_map = {
|
|
138
|
+
'round_robin': LoadBalanceStrategy.ROUND_ROBIN,
|
|
139
|
+
'least_loaded': LoadBalanceStrategy.LEAST_LOADED,
|
|
140
|
+
'random': LoadBalanceStrategy.RANDOM
|
|
141
|
+
}
|
|
142
|
+
strategy = strategy_map[ctx.obj['load_balance_strategy']]
|
|
143
|
+
|
|
144
|
+
if ctx.obj['verbose']:
|
|
145
|
+
print_info(f"Using multi-endpoint with {len(endpoints)} endpoints")
|
|
146
|
+
print_info(f"Load balance strategy: {strategy.value}")
|
|
147
|
+
|
|
148
|
+
return MultiEndpointClient(
|
|
149
|
+
endpoints=endpoints,
|
|
150
|
+
strategy=strategy,
|
|
151
|
+
timeout=ctx.obj['timeout']
|
|
152
|
+
)
|
|
153
|
+
else:
|
|
154
|
+
# Single endpoint
|
|
155
|
+
return Boltz2Client(
|
|
156
|
+
base_url=ctx.obj['base_url'],
|
|
157
|
+
api_key=ctx.obj['api_key'],
|
|
158
|
+
endpoint_type=ctx.obj['endpoint_type'],
|
|
159
|
+
timeout=ctx.obj['timeout'],
|
|
160
|
+
poll_seconds=ctx.obj['poll_seconds'],
|
|
161
|
+
console=console
|
|
162
|
+
)
|
|
117
163
|
|
|
118
164
|
|
|
119
165
|
@cli.command()
|
|
@@ -316,6 +362,8 @@ def protein(ctx, sequence: str, polymer_id: str, recycling_steps: int, sampling_
|
|
|
316
362
|
@click.option('--sampling-steps-affinity', default=200, type=click.IntRange(10, 1000), help='Sampling steps for affinity prediction (default: 200)')
|
|
317
363
|
@click.option('--diffusion-samples-affinity', default=5, type=click.IntRange(1, 10), help='Diffusion samples for affinity prediction (default: 5)')
|
|
318
364
|
@click.option('--affinity-mw-correction', is_flag=True, help='Apply molecular weight correction to affinity prediction')
|
|
365
|
+
@click.option('--msa-file', multiple=True, type=(str, click.Choice(['sto', 'a3m', 'csv', 'fasta'])),
|
|
366
|
+
help='MSA file and format (can be specified multiple times)')
|
|
319
367
|
@click.option('--output-dir', type=click.Path(), default='.', help='Directory to save output files (structure_0.cif, prediction_metadata.json). Default: current directory')
|
|
320
368
|
@click.option('--no-save', is_flag=True, help='Do not save structure files')
|
|
321
369
|
@click.pass_context
|
|
@@ -323,15 +371,17 @@ def ligand(ctx, protein_sequence: str, smiles: Optional[str], ccd: Optional[str]
|
|
|
323
371
|
protein_id: str, ligand_id: str, pocket_residues: Optional[str],
|
|
324
372
|
recycling_steps: int, sampling_steps: int, predict_affinity: bool,
|
|
325
373
|
sampling_steps_affinity: int, diffusion_samples_affinity: int,
|
|
326
|
-
affinity_mw_correction: bool,
|
|
374
|
+
affinity_mw_correction: bool, msa_file: List[Tuple[str, str]],
|
|
375
|
+
output_dir: str, no_save: bool):
|
|
327
376
|
"""
|
|
328
|
-
Predict protein-ligand complex structure.
|
|
377
|
+
Predict protein-ligand complex structure with optional MSA guidance.
|
|
329
378
|
|
|
330
379
|
PROTEIN_SEQUENCE: Protein amino acid sequence
|
|
331
380
|
|
|
332
381
|
Example:
|
|
333
382
|
boltz2 ligand "PROTEIN_SEQ" --smiles "CC(=O)OC1=CC=CC=C1C(=O)O"
|
|
334
383
|
boltz2 ligand "PROTEIN_SEQ" --ccd ASP --pocket-residues "10,15,20,25"
|
|
384
|
+
boltz2 ligand "PROTEIN_SEQ" --smiles "CC(=O)O" --msa-file alignment.a3m a3m --predict-affinity
|
|
335
385
|
"""
|
|
336
386
|
if not smiles and not ccd:
|
|
337
387
|
print_error("Must provide either --smiles or --ccd")
|
|
@@ -374,32 +424,37 @@ def ligand(ctx, protein_sequence: str, smiles: Optional[str], ccd: Optional[str]
|
|
|
374
424
|
def progress_callback(message: str):
|
|
375
425
|
progress.update(task, description=message)
|
|
376
426
|
|
|
377
|
-
#
|
|
378
|
-
|
|
379
|
-
|
|
380
|
-
|
|
381
|
-
|
|
382
|
-
|
|
427
|
+
# Prepare MSA files
|
|
428
|
+
msa_files = []
|
|
429
|
+
for file_path, format_type in msa_file:
|
|
430
|
+
if not Path(file_path).exists():
|
|
431
|
+
print_error(f"MSA file not found: {file_path}")
|
|
432
|
+
raise click.Abort()
|
|
433
|
+
msa_files.append((file_path, format_type))
|
|
383
434
|
|
|
384
|
-
|
|
385
|
-
|
|
386
|
-
smiles=smiles,
|
|
387
|
-
ccd=ccd,
|
|
388
|
-
predict_affinity=predict_affinity
|
|
389
|
-
)
|
|
435
|
+
if msa_files:
|
|
436
|
+
print_info(f"Using {len(msa_files)} MSA file(s)")
|
|
390
437
|
|
|
391
|
-
|
|
392
|
-
|
|
393
|
-
|
|
438
|
+
# Use the convenience method that handles MSA
|
|
439
|
+
result = await client.predict_protein_ligand_complex(
|
|
440
|
+
protein_sequence=protein_sequence,
|
|
441
|
+
ligand_smiles=smiles,
|
|
442
|
+
ligand_ccd=ccd,
|
|
443
|
+
protein_id=protein_id,
|
|
444
|
+
ligand_id=ligand_id,
|
|
445
|
+
pocket_residues=pocket_res_list if pocket_residues else None,
|
|
394
446
|
recycling_steps=recycling_steps,
|
|
395
447
|
sampling_steps=sampling_steps,
|
|
396
|
-
|
|
397
|
-
|
|
398
|
-
|
|
448
|
+
predict_affinity=predict_affinity,
|
|
449
|
+
sampling_steps_affinity=sampling_steps_affinity,
|
|
450
|
+
diffusion_samples_affinity=diffusion_samples_affinity,
|
|
451
|
+
affinity_mw_correction=affinity_mw_correction,
|
|
452
|
+
msa_files=msa_files if msa_files else None,
|
|
453
|
+
save_structures=not no_save,
|
|
454
|
+
output_dir=Path(output_dir),
|
|
455
|
+
progress_callback=progress_callback
|
|
399
456
|
)
|
|
400
457
|
|
|
401
|
-
result = await client.predict(request)
|
|
402
|
-
|
|
403
458
|
progress.update(task, description="Prediction completed!")
|
|
404
459
|
|
|
405
460
|
print_success(f"Complex prediction completed successfully!")
|
|
@@ -951,7 +1006,7 @@ def yaml_config(ctx, yaml_file: str, msa_dir: Optional[str], recycling_steps: in
|
|
|
951
1006
|
# Update the corresponding polymer with MSA
|
|
952
1007
|
polymer_idx = sum(1 for s in config.sequences[:i] if s.protein)
|
|
953
1008
|
if polymer_idx < len(request.polymers):
|
|
954
|
-
request.polymers[polymer_idx].msa =
|
|
1009
|
+
request.polymers[polymer_idx].msa = {"default": {format_type: msa_record}}
|
|
955
1010
|
else:
|
|
956
1011
|
print_warning(f"MSA file not found: {msa_path}")
|
|
957
1012
|
|
|
@@ -987,6 +1042,364 @@ def yaml_config(ctx, yaml_file: str, msa_dir: Optional[str], recycling_steps: in
|
|
|
987
1042
|
asyncio.run(run_yaml_prediction())
|
|
988
1043
|
|
|
989
1044
|
|
|
1045
|
+
@cli.command(name='msa-search')
|
|
1046
|
+
@click.argument('sequence')
|
|
1047
|
+
@click.option('--endpoint', default='http://your-msa-nim:8000',
|
|
1048
|
+
help='MSA Search NIM endpoint URL')
|
|
1049
|
+
@click.option('--databases', '-d', multiple=True, default=['all'],
|
|
1050
|
+
help='Databases to search (default: all)')
|
|
1051
|
+
@click.option('--max-sequences', default=500, type=int,
|
|
1052
|
+
help='Maximum sequences to return (default: 500)')
|
|
1053
|
+
@click.option('--e-value', default=0.0001, type=float,
|
|
1054
|
+
help='E-value threshold (default: 0.0001)')
|
|
1055
|
+
@click.option('--output-format', '-f',
|
|
1056
|
+
type=click.Choice(['a3m', 'fasta', 'sto']),
|
|
1057
|
+
default='a3m',
|
|
1058
|
+
help='Output format (default: a3m)')
|
|
1059
|
+
@click.option('--output', '-o', type=click.Path(), required=True,
|
|
1060
|
+
help='Output file path')
|
|
1061
|
+
@click.pass_context
|
|
1062
|
+
def msa_search_command(ctx, sequence: str, endpoint: str, databases: List[str],
|
|
1063
|
+
max_sequences: int, e_value: float,
|
|
1064
|
+
output_format: str, output: str):
|
|
1065
|
+
"""
|
|
1066
|
+
Search for MSA using GPU-accelerated MSA Search NIM.
|
|
1067
|
+
|
|
1068
|
+
Examples:
|
|
1069
|
+
|
|
1070
|
+
# Basic MSA search
|
|
1071
|
+
boltz2 msa-search "MKTVRQERLKS..." -o output.a3m
|
|
1072
|
+
|
|
1073
|
+
# Search specific databases with custom parameters
|
|
1074
|
+
boltz2 msa-search "SEQUENCE" -d uniref90 -d pdb70 --max-sequences 1000 -o output.a3m
|
|
1075
|
+
|
|
1076
|
+
# Export in different format
|
|
1077
|
+
boltz2 msa-search "SEQUENCE" -f fasta -o output.fasta
|
|
1078
|
+
"""
|
|
1079
|
+
async def run_msa_search():
|
|
1080
|
+
try:
|
|
1081
|
+
# Get client configuration
|
|
1082
|
+
config = ctx.obj
|
|
1083
|
+
client = Boltz2Client(
|
|
1084
|
+
base_url=config['base_url'],
|
|
1085
|
+
api_key=config.get('api_key'),
|
|
1086
|
+
endpoint_type=config['endpoint_type']
|
|
1087
|
+
)
|
|
1088
|
+
|
|
1089
|
+
# Configure MSA Search
|
|
1090
|
+
print_info(f"Configuring MSA Search NIM: {endpoint}")
|
|
1091
|
+
client.configure_msa_search(
|
|
1092
|
+
msa_endpoint_url=endpoint,
|
|
1093
|
+
api_key=config.get('api_key')
|
|
1094
|
+
)
|
|
1095
|
+
|
|
1096
|
+
# Show search parameters
|
|
1097
|
+
print_info("Search Parameters:")
|
|
1098
|
+
print(f" Databases: {', '.join(databases)}")
|
|
1099
|
+
print(f" Max sequences: {max_sequences}")
|
|
1100
|
+
print(f" E-value: {e_value}")
|
|
1101
|
+
print(f" Output format: {output_format}")
|
|
1102
|
+
|
|
1103
|
+
# Perform search
|
|
1104
|
+
with Progress(
|
|
1105
|
+
SpinnerColumn(),
|
|
1106
|
+
TextColumn("[progress.description]{task.description}"),
|
|
1107
|
+
BarColumn(),
|
|
1108
|
+
TimeElapsedColumn()
|
|
1109
|
+
) as progress:
|
|
1110
|
+
task = progress.add_task("Searching MSA...", total=None)
|
|
1111
|
+
|
|
1112
|
+
result_path = await client.search_msa(
|
|
1113
|
+
sequence=sequence,
|
|
1114
|
+
databases=list(databases),
|
|
1115
|
+
max_msa_sequences=max_sequences,
|
|
1116
|
+
e_value=e_value,
|
|
1117
|
+
output_format=output_format,
|
|
1118
|
+
save_path=output
|
|
1119
|
+
)
|
|
1120
|
+
|
|
1121
|
+
progress.update(task, completed=100)
|
|
1122
|
+
|
|
1123
|
+
# Show results
|
|
1124
|
+
file_size = Path(result_path).stat().st_size
|
|
1125
|
+
seq_count = Path(result_path).read_text().count('\n>')
|
|
1126
|
+
|
|
1127
|
+
print_success(f"MSA search completed!")
|
|
1128
|
+
print(f" Sequences found: {seq_count}")
|
|
1129
|
+
print(f" File size: {file_size:,} bytes")
|
|
1130
|
+
print(f" Saved to: {result_path}")
|
|
1131
|
+
|
|
1132
|
+
except Exception as e:
|
|
1133
|
+
print_error(f"MSA search failed: {e}")
|
|
1134
|
+
|
|
1135
|
+
asyncio.run(run_msa_search())
|
|
1136
|
+
|
|
1137
|
+
|
|
1138
|
+
@cli.command(name='msa-predict')
|
|
1139
|
+
@click.argument('sequence')
|
|
1140
|
+
@click.option('--endpoint', default='http://your-msa-nim:8000',
|
|
1141
|
+
help='MSA Search NIM endpoint URL')
|
|
1142
|
+
@click.option('--databases', '-d', multiple=True, default=['all'],
|
|
1143
|
+
help='Databases to search (default: all)')
|
|
1144
|
+
@click.option('--max-sequences', default=500, type=int,
|
|
1145
|
+
help='Maximum sequences for MSA (default: 500)')
|
|
1146
|
+
@click.option('--e-value', default=0.0001, type=float,
|
|
1147
|
+
help='E-value threshold (default: 0.0001)')
|
|
1148
|
+
@click.option('--recycling-steps', default=3, type=click.IntRange(1, 6),
|
|
1149
|
+
help='Number of recycling steps (default: 3)')
|
|
1150
|
+
@click.option('--sampling-steps', default=50, type=click.IntRange(10, 1000),
|
|
1151
|
+
help='Number of sampling steps (default: 50)')
|
|
1152
|
+
@click.option('--output-dir', type=click.Path(), default='.',
|
|
1153
|
+
help='Directory to save output files')
|
|
1154
|
+
@click.option('--no-save-msa', is_flag=True,
|
|
1155
|
+
help="Don't save the MSA file separately")
|
|
1156
|
+
@click.pass_context
|
|
1157
|
+
def msa_predict_command(ctx, sequence: str, endpoint: str, databases: List[str],
|
|
1158
|
+
max_sequences: int, e_value: float,
|
|
1159
|
+
recycling_steps: int, sampling_steps: int,
|
|
1160
|
+
output_dir: str, no_save_msa: bool):
|
|
1161
|
+
"""
|
|
1162
|
+
Perform MSA search and structure prediction in one step.
|
|
1163
|
+
|
|
1164
|
+
This command combines MSA search with structure prediction for enhanced results.
|
|
1165
|
+
|
|
1166
|
+
Examples:
|
|
1167
|
+
|
|
1168
|
+
# Basic MSA-guided prediction
|
|
1169
|
+
boltz2 msa-predict "MKTVRQERLKS..."
|
|
1170
|
+
|
|
1171
|
+
# Custom parameters
|
|
1172
|
+
boltz2 msa-predict "SEQUENCE" --max-sequences 1000 --recycling-steps 5
|
|
1173
|
+
|
|
1174
|
+
# Save to specific directory
|
|
1175
|
+
boltz2 msa-predict "SEQUENCE" --output-dir results/
|
|
1176
|
+
"""
|
|
1177
|
+
async def run_msa_predict():
|
|
1178
|
+
try:
|
|
1179
|
+
# Get client configuration
|
|
1180
|
+
config = ctx.obj
|
|
1181
|
+
client = Boltz2Client(
|
|
1182
|
+
base_url=config['base_url'],
|
|
1183
|
+
api_key=config.get('api_key'),
|
|
1184
|
+
endpoint_type=config['endpoint_type']
|
|
1185
|
+
)
|
|
1186
|
+
|
|
1187
|
+
# Configure MSA Search
|
|
1188
|
+
print_info(f"Configuring MSA Search NIM: {endpoint}")
|
|
1189
|
+
client.configure_msa_search(
|
|
1190
|
+
msa_endpoint_url=endpoint,
|
|
1191
|
+
api_key=config.get('api_key')
|
|
1192
|
+
)
|
|
1193
|
+
|
|
1194
|
+
# Show parameters
|
|
1195
|
+
print_info("MSA Search Parameters:")
|
|
1196
|
+
print(f" Databases: {', '.join(databases)}")
|
|
1197
|
+
print(f" Max sequences: {max_sequences}")
|
|
1198
|
+
print(f" E-value: {e_value}")
|
|
1199
|
+
|
|
1200
|
+
print_info("Prediction Parameters:")
|
|
1201
|
+
print(f" Recycling steps: {recycling_steps}")
|
|
1202
|
+
print(f" Sampling steps: {sampling_steps}")
|
|
1203
|
+
|
|
1204
|
+
# Perform MSA search + prediction
|
|
1205
|
+
with Progress(
|
|
1206
|
+
SpinnerColumn(),
|
|
1207
|
+
TextColumn("[progress.description]{task.description}"),
|
|
1208
|
+
BarColumn(),
|
|
1209
|
+
TimeElapsedColumn()
|
|
1210
|
+
) as progress:
|
|
1211
|
+
task = progress.add_task("MSA search + structure prediction...", total=None)
|
|
1212
|
+
|
|
1213
|
+
result = await client.predict_with_msa_search(
|
|
1214
|
+
sequence=sequence,
|
|
1215
|
+
databases=list(databases),
|
|
1216
|
+
max_msa_sequences=max_sequences,
|
|
1217
|
+
e_value=e_value,
|
|
1218
|
+
recycling_steps=recycling_steps,
|
|
1219
|
+
sampling_steps=sampling_steps
|
|
1220
|
+
)
|
|
1221
|
+
|
|
1222
|
+
progress.update(task, completed=100)
|
|
1223
|
+
|
|
1224
|
+
# Save results
|
|
1225
|
+
output_path = Path(output_dir)
|
|
1226
|
+
output_path.mkdir(exist_ok=True)
|
|
1227
|
+
|
|
1228
|
+
structure_file = output_path / "structure_with_msa.cif"
|
|
1229
|
+
structure_file.write_text(result.structures[0].structure)
|
|
1230
|
+
|
|
1231
|
+
confidence = result.confidence_scores[0] if result.confidence_scores else 0.0
|
|
1232
|
+
|
|
1233
|
+
print_success("Prediction completed!")
|
|
1234
|
+
print(f" Confidence score: {confidence:.3f}")
|
|
1235
|
+
print(f" Structure saved to: {structure_file}")
|
|
1236
|
+
|
|
1237
|
+
# Optionally save MSA separately
|
|
1238
|
+
if not no_save_msa:
|
|
1239
|
+
msa_file = output_path / "msa_alignment.a3m"
|
|
1240
|
+
await client.search_msa(
|
|
1241
|
+
sequence=sequence,
|
|
1242
|
+
databases=list(databases),
|
|
1243
|
+
max_msa_sequences=max_sequences,
|
|
1244
|
+
e_value=e_value,
|
|
1245
|
+
output_format='a3m',
|
|
1246
|
+
save_path=msa_file
|
|
1247
|
+
)
|
|
1248
|
+
print(f" MSA saved to: {msa_file}")
|
|
1249
|
+
|
|
1250
|
+
except Exception as e:
|
|
1251
|
+
print_error(f"MSA prediction failed: {e}")
|
|
1252
|
+
|
|
1253
|
+
asyncio.run(run_msa_predict())
|
|
1254
|
+
|
|
1255
|
+
|
|
1256
|
+
@cli.command(name='msa-ligand')
|
|
1257
|
+
@click.argument('protein_sequence')
|
|
1258
|
+
@click.option('--smiles', help='Ligand SMILES string')
|
|
1259
|
+
@click.option('--ccd', help='Ligand CCD code (alternative to SMILES)')
|
|
1260
|
+
@click.option('--endpoint', default='http://your-msa-nim:8000',
|
|
1261
|
+
help='MSA Search NIM endpoint URL')
|
|
1262
|
+
@click.option('--databases', '-d', multiple=True, default=['all'],
|
|
1263
|
+
help='Databases to search (default: all)')
|
|
1264
|
+
@click.option('--max-sequences', default=500, type=int,
|
|
1265
|
+
help='Maximum sequences for MSA (default: 500)')
|
|
1266
|
+
@click.option('--e-value', default=0.0001, type=float,
|
|
1267
|
+
help='E-value threshold (default: 0.0001)')
|
|
1268
|
+
@click.option('--predict-affinity', is_flag=True,
|
|
1269
|
+
help='Enable affinity prediction')
|
|
1270
|
+
@click.option('--sampling-steps-affinity', default=200, type=int,
|
|
1271
|
+
help='Sampling steps for affinity (default: 200)')
|
|
1272
|
+
@click.option('--diffusion-samples-affinity', default=5, type=int,
|
|
1273
|
+
help='Diffusion samples for affinity (default: 5)')
|
|
1274
|
+
@click.option('--affinity-mw-correction', is_flag=True,
|
|
1275
|
+
help='Apply MW correction to affinity')
|
|
1276
|
+
@click.option('--recycling-steps', default=3, type=click.IntRange(1, 6),
|
|
1277
|
+
help='Number of recycling steps (default: 3)')
|
|
1278
|
+
@click.option('--sampling-steps', default=50, type=click.IntRange(10, 1000),
|
|
1279
|
+
help='Number of sampling steps (default: 50)')
|
|
1280
|
+
@click.option('--output-dir', type=click.Path(), default='.',
|
|
1281
|
+
help='Directory to save output files')
|
|
1282
|
+
@click.pass_context
|
|
1283
|
+
def msa_ligand_command(ctx, protein_sequence: str, smiles: Optional[str], ccd: Optional[str],
|
|
1284
|
+
endpoint: str, databases: List[str], max_sequences: int, e_value: float,
|
|
1285
|
+
predict_affinity: bool, sampling_steps_affinity: int,
|
|
1286
|
+
diffusion_samples_affinity: int, affinity_mw_correction: bool,
|
|
1287
|
+
recycling_steps: int, sampling_steps: int, output_dir: str):
|
|
1288
|
+
"""
|
|
1289
|
+
MSA search + protein-ligand prediction with optional affinity.
|
|
1290
|
+
|
|
1291
|
+
Combines MSA search with ligand complex prediction for enhanced accuracy.
|
|
1292
|
+
|
|
1293
|
+
Examples:
|
|
1294
|
+
|
|
1295
|
+
# Basic MSA-guided ligand prediction
|
|
1296
|
+
boltz2 msa-ligand "MKTVRQERLKS..." --smiles "CC(=O)O"
|
|
1297
|
+
|
|
1298
|
+
# With affinity prediction
|
|
1299
|
+
boltz2 msa-ligand "SEQUENCE" --smiles "CC(=O)O" --predict-affinity
|
|
1300
|
+
|
|
1301
|
+
# Custom parameters
|
|
1302
|
+
boltz2 msa-ligand "SEQUENCE" --ccd ATP --max-sequences 1000 \\
|
|
1303
|
+
--predict-affinity --sampling-steps-affinity 300
|
|
1304
|
+
"""
|
|
1305
|
+
if not smiles and not ccd:
|
|
1306
|
+
print_error("Must provide either --smiles or --ccd")
|
|
1307
|
+
raise click.Abort()
|
|
1308
|
+
|
|
1309
|
+
if smiles and ccd:
|
|
1310
|
+
print_error("Provide either --smiles or --ccd, not both")
|
|
1311
|
+
raise click.Abort()
|
|
1312
|
+
|
|
1313
|
+
async def run_msa_ligand():
|
|
1314
|
+
try:
|
|
1315
|
+
# Get client configuration
|
|
1316
|
+
config = ctx.obj
|
|
1317
|
+
client = Boltz2Client(
|
|
1318
|
+
base_url=config['base_url'],
|
|
1319
|
+
api_key=config.get('api_key'),
|
|
1320
|
+
endpoint_type=config['endpoint_type']
|
|
1321
|
+
)
|
|
1322
|
+
|
|
1323
|
+
# Configure MSA Search
|
|
1324
|
+
print_info(f"Configuring MSA Search NIM: {endpoint}")
|
|
1325
|
+
client.configure_msa_search(
|
|
1326
|
+
msa_endpoint_url=endpoint,
|
|
1327
|
+
api_key=config.get('api_key')
|
|
1328
|
+
)
|
|
1329
|
+
|
|
1330
|
+
# Show parameters
|
|
1331
|
+
print_info("MSA Search Parameters:")
|
|
1332
|
+
print(f" Databases: {', '.join(databases)}")
|
|
1333
|
+
print(f" Max sequences: {max_sequences}")
|
|
1334
|
+
print(f" E-value: {e_value}")
|
|
1335
|
+
|
|
1336
|
+
print_info("Prediction Parameters:")
|
|
1337
|
+
print(f" Recycling steps: {recycling_steps}")
|
|
1338
|
+
print(f" Sampling steps: {sampling_steps}")
|
|
1339
|
+
|
|
1340
|
+
if predict_affinity:
|
|
1341
|
+
print_info("Affinity Parameters:")
|
|
1342
|
+
print(f" Sampling steps: {sampling_steps_affinity}")
|
|
1343
|
+
print(f" Diffusion samples: {diffusion_samples_affinity}")
|
|
1344
|
+
print(f" MW correction: {affinity_mw_correction}")
|
|
1345
|
+
|
|
1346
|
+
# Perform MSA search + ligand prediction
|
|
1347
|
+
with Progress(
|
|
1348
|
+
SpinnerColumn(),
|
|
1349
|
+
TextColumn("[progress.description]{task.description}"),
|
|
1350
|
+
BarColumn(),
|
|
1351
|
+
TimeElapsedColumn()
|
|
1352
|
+
) as progress:
|
|
1353
|
+
task = progress.add_task("MSA search + ligand prediction...", total=None)
|
|
1354
|
+
|
|
1355
|
+
result = await client.predict_ligand_with_msa_search(
|
|
1356
|
+
protein_sequence=protein_sequence,
|
|
1357
|
+
ligand_smiles=smiles,
|
|
1358
|
+
ligand_ccd=ccd,
|
|
1359
|
+
databases=list(databases),
|
|
1360
|
+
e_value=e_value,
|
|
1361
|
+
max_msa_sequences=max_sequences,
|
|
1362
|
+
recycling_steps=recycling_steps,
|
|
1363
|
+
sampling_steps=sampling_steps,
|
|
1364
|
+
predict_affinity=predict_affinity,
|
|
1365
|
+
sampling_steps_affinity=sampling_steps_affinity if predict_affinity else None,
|
|
1366
|
+
diffusion_samples_affinity=diffusion_samples_affinity if predict_affinity else None,
|
|
1367
|
+
affinity_mw_correction=affinity_mw_correction if predict_affinity else None,
|
|
1368
|
+
save_structures=True,
|
|
1369
|
+
output_dir=Path(output_dir)
|
|
1370
|
+
)
|
|
1371
|
+
|
|
1372
|
+
progress.update(task, completed=100)
|
|
1373
|
+
|
|
1374
|
+
# Save results
|
|
1375
|
+
output_path = Path(output_dir)
|
|
1376
|
+
|
|
1377
|
+
print_success("Prediction completed!")
|
|
1378
|
+
|
|
1379
|
+
if result.confidence_scores:
|
|
1380
|
+
confidence = result.confidence_scores[0]
|
|
1381
|
+
print(f" Confidence score: {confidence:.3f}")
|
|
1382
|
+
|
|
1383
|
+
if result.structures:
|
|
1384
|
+
structure_file = output_path / "structure_0.cif"
|
|
1385
|
+
print(f" Structure saved to: {structure_file}")
|
|
1386
|
+
|
|
1387
|
+
# Display affinity results if available
|
|
1388
|
+
if predict_affinity and result.affinities:
|
|
1389
|
+
ligand_id = "LIG" # Default ligand ID
|
|
1390
|
+
if ligand_id in result.affinities:
|
|
1391
|
+
aff = result.affinities[ligand_id]
|
|
1392
|
+
print_info("Affinity Predictions:")
|
|
1393
|
+
print(f" pIC50: {aff.affinity_pic50[0]:.3f}")
|
|
1394
|
+
print(f" IC50: {aff.affinity_ic50[0]:.3f} nM")
|
|
1395
|
+
print(f" Binding probability: {aff.affinity_probability_binary[0]:.3f}")
|
|
1396
|
+
|
|
1397
|
+
except Exception as e:
|
|
1398
|
+
print_error(f"MSA-ligand prediction failed: {e}")
|
|
1399
|
+
|
|
1400
|
+
asyncio.run(run_msa_ligand())
|
|
1401
|
+
|
|
1402
|
+
|
|
990
1403
|
@cli.command(name='screen')
|
|
991
1404
|
@click.argument('target_sequence', type=str)
|
|
992
1405
|
@click.argument('compounds_file', type=click.Path(exists=True))
|
|
@@ -1010,8 +1423,7 @@ def screen(ctx, target_sequence, compounds_file, target_name, output_dir, no_aff
|
|
|
1010
1423
|
boltz2 screen "MKTVRQERLK..." compounds.csv -o results/
|
|
1011
1424
|
boltz2 screen target.fasta library.json --pocket-residues "10,15,20,25"
|
|
1012
1425
|
"""
|
|
1013
|
-
|
|
1014
|
-
client = ctx.obj["client"]
|
|
1426
|
+
client = create_client(ctx)
|
|
1015
1427
|
|
|
1016
1428
|
# Import here to avoid circular imports
|
|
1017
1429
|
from .virtual_screening import VirtualScreening, CompoundLibrary
|