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 CHANGED
@@ -26,8 +26,9 @@ Example:
26
26
  >>> print(f"Confidence: {result.confidence_scores[0]:.3f}")
27
27
  """
28
28
 
29
- __version__ = "0.2"
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
@@ -0,0 +1,9 @@
1
+ #!/usr/bin/env python3
2
+ """
3
+ Allow running the boltz2_client module directly.
4
+ """
5
+
6
+ from .cli import cli
7
+
8
+ if __name__ == "__main__":
9
+ cli()
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, verbose: bool):
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
- print_info(f"Using {endpoint_type} endpoint: {base_url}")
100
- if endpoint_type == 'nvidia_hosted':
101
- if api_key:
102
- print_info("API key provided via command line")
103
- else:
104
- print_info("API key will be read from NVIDIA_API_KEY environment variable")
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) -> Boltz2Client:
108
- """Create a Boltz2Client from context."""
109
- return Boltz2Client(
110
- base_url=ctx.obj['base_url'],
111
- api_key=ctx.obj['api_key'],
112
- endpoint_type=ctx.obj['endpoint_type'],
113
- timeout=ctx.obj['timeout'],
114
- poll_seconds=ctx.obj['poll_seconds'],
115
- console=console
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, output_dir: str, no_save: 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
- # Create request with affinity parameters
378
- polymer = Polymer(
379
- id=protein_id,
380
- molecule_type="protein",
381
- sequence=protein_sequence
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
- ligand_obj = Ligand(
385
- id=ligand_id,
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
- request = PredictionRequest(
392
- polymers=[polymer],
393
- ligands=[ligand_obj],
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
- sampling_steps_affinity=sampling_steps_affinity if predict_affinity else None,
397
- diffusion_samples_affinity=diffusion_samples_affinity if predict_affinity else None,
398
- affinity_mw_correction=affinity_mw_correction if predict_affinity else None
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 = [msa_record]
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
- console = ctx.obj["console"]
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