boltz2-python-client 0.3.2__py3-none-any.whl → 0.3.3__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,7 +26,7 @@ Example:
26
26
  >>> print(f"Confidence: {result.confidence_scores[0]:.3f}")
27
27
  """
28
28
 
29
- __version__ = "0.3.2"
29
+ __version__ = "0.3.3"
30
30
  __author__ = "NVIDIA Corporation"
31
31
  __email__ = "bionemo-support@nvidia.com"
32
32
 
@@ -1192,6 +1192,11 @@ class A3MToCSVConverter:
1192
1192
  - Rows with the same 'key' across different chain CSVs are paired
1193
1193
  - This enables the model to identify co-evolved sequences from same organisms
1194
1194
 
1195
+ When include_unpaired=True, also includes:
1196
+ - Unpaired sequences in "block-diagonal" format where only one chain has
1197
+ a sequence and others have gaps. This maximizes MSA depth while still
1198
+ providing pairing information where available.
1199
+
1195
1200
  Args:
1196
1201
  msas: Dictionary mapping chain IDs to parsed A3MMSA objects
1197
1202
  output_path: Optional path to save the combined CSV file
@@ -1204,35 +1209,79 @@ class A3MToCSVConverter:
1204
1209
  # Find paired sequences
1205
1210
  pairs = self.pairing_strategy.find_pairs(msas)
1206
1211
 
1212
+ # Track which sequences have been paired (to find unpaired ones later)
1213
+ paired_sequence_ids: Dict[str, Set[str]] = {chain_id: set() for chain_id in chain_ids}
1214
+ for pair in pairs:
1215
+ for chain_id, seq in pair.items():
1216
+ paired_sequence_ids[chain_id].add(seq.identifier)
1217
+
1207
1218
  if self.max_pairs:
1208
1219
  pairs = pairs[:self.max_pairs]
1209
1220
 
1210
1221
  logger.info(f"Found {len(pairs)} paired sequence sets across {len(chain_ids)} chains")
1211
1222
 
1212
- # Extract query sequences
1223
+ # Extract query sequences and their lengths (for gap filling)
1213
1224
  query_sequences = {}
1225
+ query_lengths = {}
1214
1226
  for chain_id, msa in msas.items():
1215
1227
  query = msa.get_query()
1216
1228
  if query:
1217
- query_sequences[chain_id] = self._clean_sequence(query.sequence)
1229
+ clean_seq = self._clean_sequence(query.sequence)
1230
+ query_sequences[chain_id] = clean_seq
1231
+ query_lengths[chain_id] = len(clean_seq)
1232
+
1233
+ # Collect unpaired sequences if requested
1234
+ unpaired_sequences: Dict[str, List[A3MSequence]] = {chain_id: [] for chain_id in chain_ids}
1235
+ num_unpaired = 0
1236
+
1237
+ if self.include_unpaired:
1238
+ for chain_id, msa in msas.items():
1239
+ for seq in msa.sequences:
1240
+ if seq.is_query:
1241
+ continue
1242
+ if seq.identifier not in paired_sequence_ids[chain_id]:
1243
+ unpaired_sequences[chain_id].append(seq)
1244
+ num_unpaired += 1
1245
+
1246
+ logger.info(f"Found {num_unpaired} unpaired sequences to include in block-diagonal format")
1247
+
1248
+ # Calculate total rows: paired + unpaired (block-diagonal)
1249
+ current_key = len(pairs) # Start unpaired keys after paired ones
1218
1250
 
1219
1251
  # Build per-chain CSVs (for Boltz2 NIM API)
1220
1252
  csv_per_chain: Dict[str, str] = {}
1221
1253
  for chain_id in chain_ids:
1222
1254
  chain_lines = ["key,sequence"]
1255
+
1256
+ # Add paired sequences
1223
1257
  for idx, pair in enumerate(pairs, start=1):
1224
1258
  if chain_id in pair:
1225
1259
  seq = self._clean_sequence(pair[chain_id].sequence)
1226
1260
  chain_lines.append(f"{idx},{seq}")
1227
1261
  else:
1228
1262
  # Use gaps if sequence not found for this chain
1229
- ref_msa = msas[chain_id]
1230
- query = ref_msa.get_query()
1231
- if query:
1232
- gap_seq = '-' * len(self._clean_sequence(query.sequence))
1233
- chain_lines.append(f"{idx},{gap_seq}")
1263
+ gap_seq = '-' * query_lengths.get(chain_id, 0)
1264
+ chain_lines.append(f"{idx},{gap_seq}")
1265
+
1234
1266
  csv_per_chain[chain_id] = '\n'.join(chain_lines)
1235
1267
 
1268
+ # Add unpaired sequences in block-diagonal format
1269
+ if self.include_unpaired:
1270
+ unpaired_key = current_key + 1
1271
+ for source_chain_id in chain_ids:
1272
+ for unpaired_seq in unpaired_sequences[source_chain_id]:
1273
+ # Add this unpaired sequence to all chain CSVs
1274
+ for chain_id in chain_ids:
1275
+ if chain_id == source_chain_id:
1276
+ # This chain has the sequence
1277
+ seq = self._clean_sequence(unpaired_seq.sequence)
1278
+ csv_per_chain[chain_id] += f"\n{unpaired_key},{seq}"
1279
+ else:
1280
+ # Other chains get gaps
1281
+ gap_seq = '-' * query_lengths.get(chain_id, 0)
1282
+ csv_per_chain[chain_id] += f"\n{unpaired_key},{gap_seq}"
1283
+ unpaired_key += 1
1284
+
1236
1285
  # Build combined CSV (for reference/documentation)
1237
1286
  csv_lines = ["key,sequence"]
1238
1287
  for idx, pair in enumerate(pairs, start=1):
@@ -1242,15 +1291,31 @@ class A3MToCSVConverter:
1242
1291
  seq = self._clean_sequence(pair[chain_id].sequence)
1243
1292
  sequences.append(seq)
1244
1293
  else:
1245
- ref_msa = msas[chain_id]
1246
- query = ref_msa.get_query()
1247
- if query:
1248
- sequences.append('-' * len(self._clean_sequence(query.sequence)))
1294
+ sequences.append('-' * query_lengths.get(chain_id, 0))
1249
1295
  concatenated = ':'.join(sequences)
1250
1296
  csv_lines.append(f"{idx},{concatenated}")
1251
1297
 
1298
+ # Add unpaired sequences to combined CSV in block-diagonal format
1299
+ if self.include_unpaired:
1300
+ unpaired_key = current_key + 1
1301
+ for source_chain_id in chain_ids:
1302
+ for unpaired_seq in unpaired_sequences[source_chain_id]:
1303
+ sequences = []
1304
+ for chain_id in chain_ids:
1305
+ if chain_id == source_chain_id:
1306
+ seq = self._clean_sequence(unpaired_seq.sequence)
1307
+ sequences.append(seq)
1308
+ else:
1309
+ sequences.append('-' * query_lengths.get(chain_id, 0))
1310
+ concatenated = ':'.join(sequences)
1311
+ csv_lines.append(f"{unpaired_key},{concatenated}")
1312
+ unpaired_key += 1
1313
+
1252
1314
  csv_content = '\n'.join(csv_lines)
1253
1315
 
1316
+ total_rows = len(pairs) + (num_unpaired if self.include_unpaired else 0)
1317
+ logger.info(f"Total MSA rows: {total_rows} ({len(pairs)} paired + {num_unpaired if self.include_unpaired else 0} unpaired)")
1318
+
1254
1319
  # Save files if path provided
1255
1320
  output_paths_per_chain = None
1256
1321
  if output_path:
@@ -1355,6 +1420,7 @@ def convert_a3m_to_multimer_csv(
1355
1420
  output_path: Optional[Path] = None,
1356
1421
  pairing_strategy: str = "greedy",
1357
1422
  use_tax_id: Optional[bool] = None, # None = auto-detect (ColabFold default)
1423
+ include_unpaired: bool = False,
1358
1424
  max_pairs: Optional[int] = None
1359
1425
  ) -> ConversionResult:
1360
1426
  """
@@ -1374,6 +1440,10 @@ def convert_a3m_to_multimer_csv(
1374
1440
  This is how ColabFold behaves.
1375
1441
  - True: Force TaxID pairing (requires OX= or species codes in headers)
1376
1442
  - False: Force UniRef/organism ID pairing (works with standard ColabFold output)
1443
+ include_unpaired: If True, include sequences without cross-chain matches in
1444
+ "block-diagonal" format (one chain has sequence, others have gaps).
1445
+ This maximizes MSA depth while still providing pairing where available.
1446
+ Default: False (only paired sequences are included)
1377
1447
  max_pairs: Maximum number of pairs to include
1378
1448
 
1379
1449
  Returns:
@@ -1388,6 +1458,12 @@ def convert_a3m_to_multimer_csv(
1388
1458
  ... # use_tax_id=None means auto-detect (default)
1389
1459
  ... )
1390
1460
  >>> print(f"Created {result.num_pairs} paired sequences")
1461
+
1462
+ >>> # Include unpaired sequences for maximum MSA depth
1463
+ >>> result = convert_a3m_to_multimer_csv(
1464
+ ... a3m_files={'A': Path('chain_A.a3m'), 'B': Path('chain_B.a3m')},
1465
+ ... include_unpaired=True # Block-diagonal format
1466
+ ... )
1391
1467
 
1392
1468
  ColabFold Compatibility:
1393
1469
  Standard ColabFold A3M files use UniRef100 cluster IDs without TaxID information.
@@ -1417,6 +1493,7 @@ def convert_a3m_to_multimer_csv(
1417
1493
 
1418
1494
  converter = A3MToCSVConverter(
1419
1495
  pairing_strategy=strategy,
1496
+ include_unpaired=include_unpaired,
1420
1497
  max_pairs=max_pairs
1421
1498
  )
1422
1499
 
boltz2_client/cli.py CHANGED
@@ -1534,10 +1534,12 @@ def screen(ctx, target_sequence, compounds_file, target_name, output_dir, no_aff
1534
1534
  @click.option('--pairing-mode', type=click.Choice(['auto', 'taxid', 'uniref']),
1535
1535
  default='auto',
1536
1536
  help='Pairing identifier mode: auto (default, like ColabFold), taxid, or uniref')
1537
+ @click.option('--include-unpaired', is_flag=True, default=False,
1538
+ help='Include unpaired sequences in block-diagonal format (maximizes MSA depth)')
1537
1539
  @click.pass_context
1538
1540
  def convert_msa_command(ctx, a3m_files: Tuple[str, ...], chain_ids: str,
1539
1541
  output: str, max_pairs: Optional[int],
1540
- pairing_strategy: str, pairing_mode: str):
1542
+ pairing_strategy: str, pairing_mode: str, include_unpaired: bool):
1541
1543
  """
1542
1544
  Convert ColabFold A3M monomer MSA files to Boltz2 multimer CSV format.
1543
1545
 
@@ -1607,6 +1609,8 @@ def convert_msa_command(ctx, a3m_files: Tuple[str, ...], chain_ids: str,
1607
1609
  print_info(f"Output: {output}")
1608
1610
  print_info(f"Pairing strategy: {pairing_strategy}")
1609
1611
  print_info(f"Pairing mode: {pairing_mode}")
1612
+ if include_unpaired:
1613
+ print_info("Include unpaired: Yes (block-diagonal format)")
1610
1614
  if max_pairs:
1611
1615
  print_info(f"Max pairs: {max_pairs}")
1612
1616
 
@@ -1624,6 +1628,7 @@ def convert_msa_command(ctx, a3m_files: Tuple[str, ...], chain_ids: str,
1624
1628
  output_path=Path(output),
1625
1629
  pairing_strategy=pairing_strategy,
1626
1630
  use_tax_id=use_tax_id,
1631
+ include_unpaired=include_unpaired,
1627
1632
  max_pairs=max_pairs
1628
1633
  )
1629
1634
 
@@ -1666,6 +1671,8 @@ def convert_msa_command(ctx, a3m_files: Tuple[str, ...], chain_ids: str,
1666
1671
  @click.option('--pairing-mode', type=click.Choice(['auto', 'taxid', 'uniref']),
1667
1672
  default='auto',
1668
1673
  help='Pairing identifier mode: auto (default), taxid, or uniref')
1674
+ @click.option('--include-unpaired', is_flag=True, default=False,
1675
+ help='Include unpaired sequences in block-diagonal format (maximizes MSA depth)')
1669
1676
  @click.option('--recycling-steps', type=int, default=3,
1670
1677
  help='Number of recycling steps (default: 3)')
1671
1678
  @click.option('--sampling-steps', type=int, default=200,
@@ -1675,7 +1682,7 @@ def convert_msa_command(ctx, a3m_files: Tuple[str, ...], chain_ids: str,
1675
1682
  @click.pass_context
1676
1683
  def multimer_msa_command(ctx, a3m_files: Tuple[str, ...], chain_ids: str,
1677
1684
  output: Optional[str], save_csv: bool, save_all: bool,
1678
- max_pairs: Optional[int], pairing_mode: str,
1685
+ max_pairs: Optional[int], pairing_mode: str, include_unpaired: bool,
1679
1686
  recycling_steps: int, sampling_steps: int,
1680
1687
  diffusion_samples: int):
1681
1688
  """
@@ -1775,10 +1782,12 @@ def multimer_msa_command(ctx, a3m_files: Tuple[str, ...], chain_ids: str,
1775
1782
  a3m_files=a3m_file_dict,
1776
1783
  pairing_strategy='greedy',
1777
1784
  use_tax_id=use_tax_id,
1785
+ include_unpaired=include_unpaired,
1778
1786
  max_pairs=max_pairs
1779
1787
  )
1780
1788
 
1781
- progress.update(task, description=f"✓ Paired {result.num_pairs} sequences")
1789
+ unpaired_msg = " (+ unpaired)" if include_unpaired else ""
1790
+ progress.update(task, description=f"✓ Paired {result.num_pairs} sequences{unpaired_msg}")
1782
1791
 
1783
1792
  print_info(f"Paired sequences: {result.num_pairs}")
1784
1793
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: boltz2-python-client
3
- Version: 0.3.2
3
+ Version: 0.3.3
4
4
  Summary: Python client for Boltz-2 protein structure prediction API with covalent complex and multi-endpoint support
5
5
  Author-email: NVIDIA Corporation <bionemo-support@nvidia.com>
6
6
  Maintainer-email: NVIDIA Corporation <bionemo-support@nvidia.com>
@@ -1,7 +1,7 @@
1
- boltz2_client/__init__.py,sha256=gQfHL4fPYBPoaGYpCGuy-WAZClzYZsSiRqDNEuEN2IA,5321
1
+ boltz2_client/__init__.py,sha256=YAjmLUxO4KE2a8wVJ4WA1tXFEDRnNalN_yOf6gNb7rk,5321
2
2
  boltz2_client/__main__.py,sha256=LamY2ky-DaWf9mGqDZljQwotWICRNeSwq9Y__erB3bQ,140
3
- boltz2_client/a3m_to_csv_converter.py,sha256=_7b7XHcfu-RttrA40SzsqBFlVw8-KHvEUkryQtgEpkk,60644
4
- boltz2_client/cli.py,sha256=nnJMv0aK_LmiD0vxbZKqZiG64jci5wlueoPD1BzuI2A,87321
3
+ boltz2_client/a3m_to_csv_converter.py,sha256=4ba0P_7Aqg80Hc4rKZZ_zegpIF-eyjmXT6c_3yo8CdM,64705
4
+ boltz2_client/cli.py,sha256=bo6aYI0pEtoeRPW6zBpeymxum4vxdJJBE5RSps042Rg,87970
5
5
  boltz2_client/client.py,sha256=FDz1X0lnB5B0Oi0anHVFYY-VZpVWjyieyX9APmapKOg,57125
6
6
  boltz2_client/exceptions.py,sha256=x-5MvZfsN1sLVJjiy8DMNfdmzJDxTosl2QGYKnrhz5k,6725
7
7
  boltz2_client/models.py,sha256=XHFTq8h0pUzK5FmOxvBm9aqZNosF75r7CMdEvtV2A50,20171
@@ -12,9 +12,9 @@ boltz2_client/utils.py,sha256=myk4O1ZiJmUvQ-1gQfj_SiDqh3wBINAsLFxH5P5X-cQ,12193
12
12
  boltz2_client/virtual_screening.py,sha256=_ykwsjg8to8MAGZGrsaOmiS3Bt9mH5DF_RUxl2Dy5ZU,23857
13
13
  boltz2_client/data/__init__.py,sha256=Qy01JhGtCtEJVIMA16fks-ZbL5r0z5JOqHri9bjkWVY,112
14
14
  boltz2_client/data/speclist.txt,sha256=5e0no6pTa5HsgErt2fMavdzZs8NZ_iZzBHEfiWKQ1lA,2438338
15
- boltz2_python_client-0.3.2.dist-info/licenses/LICENSE,sha256=7sYV_pvGtWzybTdOcFXLDPRZL2VolpIl6nt8akVEO94,1111
16
- boltz2_python_client-0.3.2.dist-info/METADATA,sha256=zdIzySJt-zNVLiEWf-yTeSOtkyEdjjnuc4bqnixpu9M,23920
17
- boltz2_python_client-0.3.2.dist-info/WHEEL,sha256=_zCd3N1l69ArxyTb8rzEoP9TpbYXkqRFSNOD5OuxnTs,91
18
- boltz2_python_client-0.3.2.dist-info/entry_points.txt,sha256=3whjBtXqltfepvt46DSJNpeUGCZGLRFYnccMEVXMCes,49
19
- boltz2_python_client-0.3.2.dist-info/top_level.txt,sha256=INM1MYyL_h2nQkTFjZE7AmrgX4ZePTr4uw4n7Sag3ao,14
20
- boltz2_python_client-0.3.2.dist-info/RECORD,,
15
+ boltz2_python_client-0.3.3.dist-info/licenses/LICENSE,sha256=7sYV_pvGtWzybTdOcFXLDPRZL2VolpIl6nt8akVEO94,1111
16
+ boltz2_python_client-0.3.3.dist-info/METADATA,sha256=dTP4OStXEXbsiZNDbNN7ZA2fuebBC_QyR9g4-E75ZT8,23920
17
+ boltz2_python_client-0.3.3.dist-info/WHEEL,sha256=wUyA8OaulRlbfwMtmQsvNngGrxQHAvkKcvRmdizlJi0,92
18
+ boltz2_python_client-0.3.3.dist-info/entry_points.txt,sha256=3whjBtXqltfepvt46DSJNpeUGCZGLRFYnccMEVXMCes,49
19
+ boltz2_python_client-0.3.3.dist-info/top_level.txt,sha256=INM1MYyL_h2nQkTFjZE7AmrgX4ZePTr4uw4n7Sag3ao,14
20
+ boltz2_python_client-0.3.3.dist-info/RECORD,,
@@ -1,5 +1,5 @@
1
1
  Wheel-Version: 1.0
2
- Generator: setuptools (80.9.0)
2
+ Generator: setuptools (80.10.2)
3
3
  Root-Is-Purelib: true
4
4
  Tag: py3-none-any
5
5