manimol 0.2.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.
Files changed (67) hide show
  1. args_parse.py +593 -0
  2. candidate_corrector.py +161 -0
  3. candidate_scorer.py +359 -0
  4. dataset/__init__.py +1 -0
  5. dataset/drugdataset.py +644 -0
  6. dataset/manifold.py +237 -0
  7. dataset/sharded_drugdataset.py +124 -0
  8. dataset/smiles2graph.py +217 -0
  9. exputils.py +289 -0
  10. infer1.py +5430 -0
  11. infer_mean_dispersion_torsion.py +1112 -0
  12. infer_torsion.py +752 -0
  13. infer_torsion5.py +1021 -0
  14. manimol/__init__.py +18 -0
  15. manimol/__main__.py +6 -0
  16. manimol/cli.py +36 -0
  17. manimol/full_cli.py +22 -0
  18. manimol/priors.py +60 -0
  19. manimol/runtime.py +116 -0
  20. manimol/selection.py +85 -0
  21. manimol/torsion.py +22 -0
  22. manimol-0.2.0.dist-info/METADATA +185 -0
  23. manimol-0.2.0.dist-info/RECORD +67 -0
  24. manimol-0.2.0.dist-info/WHEEL +5 -0
  25. manimol-0.2.0.dist-info/entry_points.txt +3 -0
  26. manimol-0.2.0.dist-info/licenses/LICENSE +21 -0
  27. manimol-0.2.0.dist-info/top_level.txt +21 -0
  28. manimol_stage1.py +206 -0
  29. mixture_flow/src/__init__.py +1 -0
  30. mixture_flow/src/_path.py +7 -0
  31. mixture_flow/src/infer28_mixture_torsion_flow.py +1093 -0
  32. mixture_flow/src/mixture_torsion_flow.py +291 -0
  33. mixture_flow/src/summarize_mixture_debug.py +80 -0
  34. mixture_flow/src/train28_mixture_torsion_flow.py +583 -0
  35. mixture_torsion_flow.py +7 -0
  36. models/ manifold_learning.py +86 -0
  37. models/__init__.py +8 -0
  38. models/dist.py +273 -0
  39. models/dist2coords.py +23 -0
  40. models/egnn.py +107 -0
  41. models/experimental_backbones.py +364 -0
  42. models/gnnconv.py +512 -0
  43. models/kernel_inversion.py +141 -0
  44. models/kernels.py +427 -0
  45. models/losses.py +154 -0
  46. models/mean_dispersion_bridge.py +119 -0
  47. models/mean_dispersion_torsion_context.py +155 -0
  48. models/model.py +73 -0
  49. models/model0.py +2375 -0
  50. models/vis.py +128 -0
  51. pgraph_controls.py +107 -0
  52. scripts/__init__.py +2 -0
  53. scripts/reference_free_geometry.py +112 -0
  54. torsion_manifold.py +281 -0
  55. train2.py +1631 -0
  56. train5.py +415 -0
  57. trainer.py +1038 -0
  58. utils/__init__.py +7 -0
  59. utils/checkpoint.py +40 -0
  60. utils/device.py +17 -0
  61. utils/helpers.py +98 -0
  62. utils/kabsch.py +29 -0
  63. utils/lookup_table.py +164 -0
  64. utils/metrics.py +50 -0
  65. utils/optuna.py +111 -0
  66. utils/save_mol.py +366 -0
  67. utils/util.py +36 -0
args_parse.py ADDED
@@ -0,0 +1,593 @@
1
+ import argparse
2
+
3
+
4
+ def args_parser():
5
+ parser = argparse.ArgumentParser(
6
+ description=(
7
+ "Probability-manifold molecular conformation generation: "
8
+ "Stage 1 manifold learning and Stage 2 multimode conformer generation."
9
+ )
10
+ )
11
+
12
+ # =========================================================
13
+ # Experiment
14
+ # =========================================================
15
+ parser.add_argument("--exp_name", type=str, default="run")
16
+ parser.add_argument("--dump_path", type=str, default="dump/")
17
+ parser.add_argument("--exp_id", type=str, default="")
18
+ parser.add_argument(
19
+ "--fixed_dump_name",
20
+ action="store_true",
21
+ help="Store runs under dump/exp_name/exp_id instead of dump/MMDD-exp_name/exp_id.",
22
+ )
23
+ parser.add_argument("--gpu", type=str, default="0")
24
+ parser.add_argument("--auto_device_config", dest="auto_device_config", action="store_true", default=True)
25
+ parser.add_argument("--no_auto_device_config", dest="auto_device_config", action="store_false")
26
+ parser.add_argument("--cpu_bs", type=int, default=32)
27
+ parser.add_argument("--cuda_bs", type=int, default=96)
28
+ parser.add_argument("--num_workers", type=int, default=None)
29
+ parser.add_argument("--cpu_num_workers", type=int, default=0)
30
+ parser.add_argument("--cuda_num_workers", type=int, default=0)
31
+ parser.add_argument("--prefetch_factor", type=int, default=None)
32
+ parser.add_argument("--cuda_prefetch_factor", type=int, default=2)
33
+ parser.add_argument("--pin_memory", dest="pin_memory", action="store_true", default=None)
34
+ parser.add_argument("--no_pin_memory", dest="pin_memory", action="store_false")
35
+ parser.add_argument("--random_seed", type=int, default=0)
36
+ parser.add_argument("--train_mol_limit", type=int, default=0)
37
+ parser.add_argument("--valid_mol_limit", type=int, default=0)
38
+ parser.add_argument("--test_mol_limit", type=int, default=0)
39
+
40
+ # checkpoint / resume
41
+ parser.add_argument("--load_path", type=str, default=None)
42
+ parser.add_argument("--checkpoint_path", type=str, default=None)
43
+ parser.add_argument(
44
+ "--resume_progress_checkpoint",
45
+ type=str,
46
+ default=None,
47
+ help=(
48
+ "Resume a MANIMOL training progress checkpoint, restoring model, "
49
+ "optimizer, epoch, batch position, and global step."
50
+ ),
51
+ )
52
+ parser.add_argument(
53
+ "--resume_next_epoch",
54
+ action="store_true",
55
+ default=False,
56
+ help="Resume from the next epoch without iterating over the unfinished part of a saved epoch.",
57
+ )
58
+ parser.add_argument("--checkpoint_interval_steps", type=int, default=1000)
59
+
60
+ # =========================================================
61
+ # Dataset
62
+ # =========================================================
63
+ parser.add_argument("--data_root", type=str, default="data")
64
+ parser.add_argument("--config_path", type=str, default="configs")
65
+ parser.add_argument("--dataset", type=str, default="QM9")
66
+ parser.add_argument(
67
+ "--raw_prefix",
68
+ type=str,
69
+ default="converted",
70
+ choices=["converted", "geom_qm9", "geom_qm9_full", "geom_drugs", "geom_drugs75k_subset"],
71
+ help="Raw data file set under data_root. geom_drugs75k_subset uses geom_drugs75k_subset_train/val/test.",
72
+ )
73
+ parser.add_argument("--train_raw_filename", type=str, default=None)
74
+ parser.add_argument("--val_raw_filename", type=str, default=None)
75
+ parser.add_argument("--test_raw_filename", type=str, default=None)
76
+
77
+ # 0 opt-in auto-detects the train/validation conformer capacity.
78
+ parser.add_argument("--max_confs", type=int, default=5)
79
+ parser.add_argument("--full_test_conformers", action="store_true", default=False,
80
+ help="Opt-in: retain every test conformer in a separate v2 cache.")
81
+ parser.add_argument("--sharded_data_root", type=str, default="",
82
+ help="Optional isolated tensor-only shard cache; legacy InMemoryDataset remains default.")
83
+ parser.add_argument("--shard_cache_size", type=int, default=2)
84
+ parser.add_argument(
85
+ "--stage2_assignment_cache",
86
+ type=str,
87
+ default="",
88
+ help="Optional directory containing fixed per-split Stage-II cluster assignments.",
89
+ )
90
+ parser.add_argument(
91
+ "--distributed_stage2",
92
+ action="store_true",
93
+ default=False,
94
+ help="Enable synchronous multi-GPU Stage-II training under torchrun.",
95
+ )
96
+
97
+ # =========================================================
98
+ # Model: GNN encoder
99
+ # =========================================================
100
+ parser.add_argument("--emb_dim", type=int, default=128)
101
+ parser.add_argument("--layer", type=int, default=4)
102
+ parser.add_argument("--dropout", type=float, default=0.5)
103
+
104
+ parser.add_argument(
105
+ "--gnn_type",
106
+ type=str,
107
+ default="gin",
108
+ choices=["gcn", "gin"],
109
+ )
110
+
111
+ parser.add_argument(
112
+ "--pooling_type",
113
+ type=str,
114
+ default="mean",
115
+ choices=["mean", "add", "max"],
116
+ )
117
+
118
+ # Input feature dimensions. These must match dataset / gnnconv.py.
119
+ parser.add_argument("--node_dim", type=int, default=39)
120
+ parser.add_argument("--edge_dim", type=int, default=13)
121
+
122
+ # =========================================================
123
+ # Model: probability map / MLP hidden dims
124
+ # =========================================================
125
+ parser.add_argument("--mlp_hidden", type=int, default=128)
126
+ parser.add_argument("--mlp_layer", type=int, default=4)
127
+
128
+ # =========================================================
129
+ # Model: probability mode generation
130
+ # =========================================================
131
+ # Number of generated probability modes/conformers during multimode training/inference.
132
+ parser.add_argument("--num_modes", type=int, default=10)
133
+
134
+ # Latent code dimension for P^(k) = f(G, z_k).
135
+ # Kept as mode_dim to avoid CVAE wording.
136
+ parser.add_argument("--mode_dim", type=int, default=32)
137
+
138
+ # Compatibility alias only. train.py should map latent_dim -> mode_dim if needed.
139
+ parser.add_argument("--latent_dim", type=int, default=None)
140
+
141
+ # =========================================================
142
+ # Model: coordinate generation / refinement
143
+ # =========================================================
144
+ parser.add_argument("--coarse_hidden", type=int, default=128)
145
+ parser.add_argument("--refine_hidden", type=int, default=128)
146
+ parser.add_argument("--num_refine_steps", type=int, default=4)
147
+
148
+ # Compatibility alias only. train.py should map refine_steps -> num_refine_steps if needed.
149
+ parser.add_argument("--refine_steps", type=int, default=None)
150
+
151
+ # =========================================================
152
+ # Training
153
+ # =========================================================
154
+ parser.add_argument("--lr", type=float, default=1e-4)
155
+ parser.add_argument("--bs", type=int, default=32)
156
+ parser.add_argument("--epoch", type=int, default=300)
157
+ parser.add_argument("--eta_min", type=float, default=1e-6)
158
+ parser.add_argument("--patience", type=int, default=20)
159
+ parser.add_argument("--lambda_torsion_supervision", type=float, default=0.0)
160
+
161
+ parser.add_argument(
162
+ "--metric",
163
+ type=str,
164
+ default="RMSD",
165
+ choices=["CE", "RMSD"],
166
+ )
167
+
168
+ # Default True, with explicit opt-out.
169
+ parser.add_argument("--early_stop", dest="early_stop", action="store_true", default=True)
170
+ parser.add_argument("--no_early_stop", dest="early_stop", action="store_false")
171
+
172
+ parser.add_argument(
173
+ "--train_stage",
174
+ type=str,
175
+ default="gnn",
176
+ choices=[
177
+ "gnn",
178
+ "multimode",
179
+ ],
180
+ )
181
+
182
+ # Compatibility. Current code always uses teacher-style supervision.
183
+ # Kept to avoid breaking old commands, but should not affect logic.
184
+ parser.add_argument("--train_model", type=str, default="teacher")
185
+
186
+ # =========================================================
187
+ # Losses: Stage 1 probability manifold alignment
188
+ # =========================================================
189
+ # Main CE weight for CE(P_global, Q_global).
190
+ parser.add_argument("--lambda_global", type=float, default=1.0)
191
+ parser.add_argument(
192
+ "--mean_dispersion_prior",
193
+ action="store_true",
194
+ help=(
195
+ "Train the optional two-head Stage-I prior: a mean pairwise proximity map "
196
+ "and an ensemble-derived pairwise dispersion map."
197
+ ),
198
+ )
199
+ parser.add_argument(
200
+ "--lambda_stage1_dispersion",
201
+ type=float,
202
+ default=1.0,
203
+ help="Weight of the pairwise dispersion regression loss in the mean--dispersion Stage-I objective.",
204
+ )
205
+ parser.add_argument(
206
+ "--stage1_target_strategy",
207
+ type=str,
208
+ default="conformer",
209
+ choices=["fusion", "lowest_energy", "first", "conformer"],
210
+ help=(
211
+ "How Stage 1 builds the teacher probability graph. "
212
+ "fusion keeps the old learnable weighted fusion; lowest_energy uses the lowest-energy conformer; "
213
+ "first uses the first conformer; conformer averages per-conformer CE without constructing a fused Q_global."
214
+ ),
215
+ )
216
+ parser.add_argument(
217
+ "--stage1_pair_relevance_weight",
218
+ type=float,
219
+ default=0.0,
220
+ help=(
221
+ "Optional isolated Route-2 extension: upweight atom pairs whose teacher proximity varies "
222
+ "across the reference conformer ensemble. Zero preserves the legacy Stage-I objective."
223
+ ),
224
+ )
225
+ parser.add_argument(
226
+ "--stage1_pair_relevance_clip",
227
+ type=float,
228
+ default=4.0,
229
+ help="Maximum normalized conformer-variation weight before applying stage1_pair_relevance_weight.",
230
+ )
231
+ parser.add_argument(
232
+ "--stage1_nonlocal_heavy_only",
233
+ action="store_true",
234
+ help=(
235
+ "Restrict Stage-I probability supervision to heavy-atom pairs separated by at least "
236
+ "three bonds. This removes bonded and angle-local pairs that cannot distinguish conformers."
237
+ ),
238
+ )
239
+ parser.add_argument(
240
+ "--functional_teacher_mode",
241
+ type=str,
242
+ default="geometry",
243
+ choices=["geometry", "functional", "dual"],
244
+ help=(
245
+ "Optional Stage-I teacher extension. 'geometry' uses the original proximity teacher; "
246
+ "'functional' uses a geometry-preserving teacher enriched for intramolecular functional pairs; "
247
+ "'dual' jointly aligns to both teachers. This never uses protein structures."
248
+ ),
249
+ )
250
+ parser.add_argument(
251
+ "--functional_teacher_alpha",
252
+ type=float,
253
+ default=0.80,
254
+ help="Geometry retention in the functional-aware teacher Q_mix = alpha*Q_geom + (1-alpha)*Q_func.",
255
+ )
256
+ parser.add_argument(
257
+ "--lambda_stage1_geometry",
258
+ type=float,
259
+ default=1.0,
260
+ help="Geometry-teacher BCE weight when --functional_teacher_mode=dual.",
261
+ )
262
+ parser.add_argument(
263
+ "--lambda_stage1_functional",
264
+ type=float,
265
+ default=0.25,
266
+ help="Functional-aware teacher BCE weight when --functional_teacher_mode=dual.",
267
+ )
268
+ parser.add_argument(
269
+ "--functional_hbond_weight",
270
+ type=float,
271
+ default=0.35,
272
+ help="Relative teacher enrichment for intramolecular donor--acceptor atom pairs.",
273
+ )
274
+ parser.add_argument(
275
+ "--functional_aromatic_weight",
276
+ type=float,
277
+ default=0.20,
278
+ help="Relative teacher enrichment for aromatic atom pairs.",
279
+ )
280
+ parser.add_argument(
281
+ "--functional_hydrophobic_weight",
282
+ type=float,
283
+ default=0.10,
284
+ help="Relative teacher enrichment for hydrophobic atom pairs.",
285
+ )
286
+ parser.add_argument(
287
+ "--stage1_kernel",
288
+ type=str,
289
+ default="umap",
290
+ choices=[
291
+ "umap",
292
+ "gaussian",
293
+ "laplacian",
294
+ "student",
295
+ "cauchy",
296
+ "rational_quadratic",
297
+ "inverse_multiquadric",
298
+ "matern32",
299
+ "matern52",
300
+ "triweight",
301
+ "logistic",
302
+ ],
303
+ help=(
304
+ "Kernel used to build teacher Q(Y) in Stage 1. "
305
+ "umap uses learnable a,b; gaussian/laplacian/cauchy/matern/triweight/logistic use learnable bandwidth from a; "
306
+ "student and rational_quadratic use b as a learnable tail/shape parameter."
307
+ ),
308
+ )
309
+ parser.add_argument(
310
+ "--stage1_fixed_kernel_a",
311
+ type=float,
312
+ default=0.0,
313
+ help="Positive value fixes kernel parameter a during Stage-I and inference; zero keeps it learnable.",
314
+ )
315
+ parser.add_argument(
316
+ "--stage1_fixed_kernel_b",
317
+ type=float,
318
+ default=0.0,
319
+ help="Positive value fixes kernel parameter b during Stage-I and inference; zero keeps it learnable.",
320
+ )
321
+
322
+ # Learnable a,b regularization. Keep this: a,b are part of teacher Q construction.
323
+ parser.add_argument("--lambda_ab", type=float, default=1e-4)
324
+
325
+ # =========================================================
326
+ # Losses: Stage 2 coarse coordinate generation
327
+ # =========================================================
328
+ parser.add_argument("--lambda_coarse_rmsd", type=float, default=1.0)
329
+
330
+ parser.add_argument("--lambda_refine_rmsd", type=float, default=1.0)
331
+ parser.add_argument(
332
+ "--lambda_bond",
333
+ type=float,
334
+ default=0.0,
335
+ help="Stage 2 bonded-distance matching loss weight, using each reference conformer's bond distances.",
336
+ )
337
+ parser.add_argument(
338
+ "--lambda_angle",
339
+ type=float,
340
+ default=0.0,
341
+ help="Stage 2 lightweight angle regularization weight.",
342
+ )
343
+ parser.add_argument(
344
+ "--lambda_pair_dist",
345
+ type=float,
346
+ default=0.0,
347
+ help="Stage 2 all-pair heavy-atom distance matching weight for refined coordinates.",
348
+ )
349
+ parser.add_argument(
350
+ "--lambda_coarse_pair_dist",
351
+ type=float,
352
+ default=0.0,
353
+ help="Stage 2 all-pair heavy-atom distance matching weight for coarse decoder coordinates.",
354
+ )
355
+ parser.add_argument(
356
+ "--pair_dist_max_pairs",
357
+ type=int,
358
+ default=2048,
359
+ help="Maximum atom pairs used by the pairwise distance loss per conformer.",
360
+ )
361
+ parser.add_argument(
362
+ "--pair_dist_min_ref",
363
+ type=float,
364
+ default=0.0,
365
+ help="Minimum reference pair distance included in the pairwise distance loss.",
366
+ )
367
+ parser.add_argument(
368
+ "--pair_dist_max_ref",
369
+ type=float,
370
+ default=12.0,
371
+ help="Maximum reference pair distance included in the pairwise distance loss.",
372
+ )
373
+ parser.add_argument(
374
+ "--lambda_clash",
375
+ type=float,
376
+ default=0.0,
377
+ help="Stage 2 non-bonded steric clash penalty weight.",
378
+ )
379
+ parser.add_argument(
380
+ "--lambda_steric",
381
+ type=float,
382
+ default=0.0,
383
+ help="Backward-compatible alias for --lambda_clash.",
384
+ )
385
+
386
+ # Backward-compatible alias.
387
+ parser.add_argument("--lambda_rmsd", type=float, default=None)
388
+
389
+ # =========================================================
390
+ # Losses: Stage 2 manifold preservation
391
+ # =========================================================
392
+ parser.add_argument("--lambda_ce_manifold", type=float, default=0.1)
393
+ parser.add_argument(
394
+ "--lambda_set_recall",
395
+ type=float,
396
+ default=0.0,
397
+ help="Stage 2 set-level recall loss weight: each reference conformer should be covered by some generated mode.",
398
+ )
399
+ parser.add_argument(
400
+ "--lambda_set_precision",
401
+ type=float,
402
+ default=0.0,
403
+ help="Stage 2 set-level precision loss weight: each generated mode should stay near some reference conformer.",
404
+ )
405
+ parser.add_argument(
406
+ "--set_loss_clamp",
407
+ type=float,
408
+ default=10.0,
409
+ help="Clamp value for RMSD entries used by set-level coverage loss.",
410
+ )
411
+ parser.add_argument(
412
+ "--detach_refiner_input",
413
+ dest="detach_refiner_input",
414
+ action="store_true",
415
+ default=True,
416
+ help="Detach coarse coordinates before refiner, preserving the legacy Stage 2 behavior.",
417
+ )
418
+ parser.add_argument(
419
+ "--no_detach_refiner_input",
420
+ dest="detach_refiner_input",
421
+ action="store_false",
422
+ help="Allow refiner losses to backpropagate into the coarse coordinate decoder.",
423
+ )
424
+ parser.add_argument(
425
+ "--lambda_mode_diversity",
426
+ type=float,
427
+ default=0.0,
428
+ help="Stage 2 mode diversity weight. Encourages generated modes for the same molecule to spread without changing P_global.",
429
+ )
430
+ parser.add_argument(
431
+ "--use_end2end_pglobal",
432
+ action="store_true",
433
+ help="Stage 2: unfreeze GNN/P_global branch and fine-tune it with a smaller learning rate.",
434
+ )
435
+ parser.add_argument(
436
+ "--pglobal_lr",
437
+ type=float,
438
+ default=3e-5,
439
+ help="Learning rate for GNN/P_global branch when --use_end2end_pglobal is enabled.",
440
+ )
441
+ parser.add_argument(
442
+ "--use_mode_specific_p",
443
+ action="store_true",
444
+ help=(
445
+ "Stage 2: generate a mode-specific probability graph P_k for each latent mode. "
446
+ "This keeps Stage 1 P_global frozen and learns P_k = f(P_global, z_k) in the decoder stage."
447
+ ),
448
+ )
449
+ parser.add_argument(
450
+ "--lambda_mode_p_ce",
451
+ type=float,
452
+ default=0.0,
453
+ help="Stage 2 weight for CE(P_k, Q_i) when a generated mode is matched to a reference conformer.",
454
+ )
455
+ parser.add_argument(
456
+ "--mode_p_delta_scale",
457
+ type=float,
458
+ default=1.0,
459
+ help="Maximum logit-space residual scale used by the mode-specific probability generator.",
460
+ )
461
+ parser.add_argument(
462
+ "--mode_assignment_strategy",
463
+ type=str,
464
+ default="min",
465
+ choices=["min", "conf_id", "cluster"],
466
+ help=(
467
+ "Stage 2 mode-reference assignment. min keeps the old min-over-modes objective; "
468
+ "conf_id assigns each reference conformer to a stable mode by conf_id modulo num_modes; "
469
+ "cluster groups reference conformers by heavy-atom RMSD with farthest-point anchors."
470
+ ),
471
+ )
472
+ parser.add_argument(
473
+ "--deterministic_modes",
474
+ action="store_true",
475
+ help="Use deterministic one-hot mode codes z_k instead of random z during Stage 2 training and inference.",
476
+ )
477
+ parser.add_argument(
478
+ "--use_mode_pair_attention",
479
+ action="store_true",
480
+ help="Stage 2: enable mode-conditioned atom-pair attention in the coordinate decoder.",
481
+ )
482
+ parser.add_argument(
483
+ "--mode_attention_heads",
484
+ type=int,
485
+ default=4,
486
+ help="Number of heads for mode-conditioned pair attention.",
487
+ )
488
+ parser.add_argument(
489
+ "--mode_attention_dropout",
490
+ type=float,
491
+ default=0.0,
492
+ help="Dropout used inside mode-conditioned pair attention.",
493
+ )
494
+ parser.add_argument(
495
+ "--mode_diversity_margin",
496
+ type=float,
497
+ default=0.6,
498
+ help="Fallback heavy-atom RMSD margin for mode diversity when reference diversity is unavailable.",
499
+ )
500
+ parser.add_argument(
501
+ "--mode_diversity_ref_scale",
502
+ type=float,
503
+ default=0.7,
504
+ help="Scale applied to average reference-reference RMSD to set the diversity margin.",
505
+ )
506
+ parser.add_argument(
507
+ "--mode_diversity_max_pairs",
508
+ type=int,
509
+ default=15,
510
+ help="Maximum generated mode pairs used by the diversity loss per molecule.",
511
+ )
512
+ parser.add_argument(
513
+ "--lambda_signed_torsion",
514
+ type=float,
515
+ default=0.0,
516
+ help="Stage 2 signed torsion supervision weight. Penalizes mirror-like torsion flips.",
517
+ )
518
+ parser.add_argument(
519
+ "--use_torsion_residual",
520
+ action="store_true",
521
+ help="Stage 2: predict and apply learned residual rotations on rotatable torsions after coordinate refinement.",
522
+ )
523
+ parser.add_argument(
524
+ "--lambda_torsion_residual_rmsd",
525
+ type=float,
526
+ default=0.0,
527
+ help="Additional RMSD weight for the torsion-residual refined coordinates.",
528
+ )
529
+ parser.add_argument(
530
+ "--lambda_torsion_delta",
531
+ type=float,
532
+ default=0.0,
533
+ help="Supervise predicted torsion residuals against reference torsion angle deltas.",
534
+ )
535
+ parser.add_argument(
536
+ "--torsion_residual_max",
537
+ type=int,
538
+ default=32,
539
+ help="Maximum number of torsions updated by the learned torsion residual module.",
540
+ )
541
+ parser.add_argument(
542
+ "--torsion_residual_scale",
543
+ type=float,
544
+ default=1.0,
545
+ help="Maximum absolute learned torsion residual in radians.",
546
+ )
547
+ parser.add_argument(
548
+ "--lambda_chiral_volume",
549
+ type=float,
550
+ default=0.0,
551
+ help="Stage 2 signed local volume supervision weight. Penalizes mirror-like local handedness flips.",
552
+ )
553
+ parser.add_argument("--signed_torsion_max", type=int, default=64)
554
+ parser.add_argument("--chiral_volume_max_centers", type=int, default=64)
555
+ parser.add_argument("--chiral_volume_ref_abs_min", type=float, default=0.05)
556
+
557
+ # =========================================================
558
+ # Evaluation / inference helpers
559
+ # =========================================================
560
+ parser.add_argument("--eval_num_samples", type=int, default=100)
561
+
562
+ # =========================================================
563
+ # Visualization
564
+ # =========================================================
565
+ parser.add_argument("--smiles", type=str, default=None)
566
+ parser.add_argument("--get_image", action="store_true")
567
+
568
+ args = parser.parse_args()
569
+
570
+ # =========================================================
571
+ # Compatibility normalization
572
+ # =========================================================
573
+ if args.latent_dim is not None:
574
+ args.mode_dim = int(args.latent_dim)
575
+ else:
576
+ args.latent_dim = int(args.mode_dim)
577
+
578
+ if args.refine_steps is not None:
579
+ args.num_refine_steps = int(args.refine_steps)
580
+ else:
581
+ args.refine_steps = int(args.num_refine_steps)
582
+
583
+ if args.lambda_rmsd is not None:
584
+ args.lambda_refine_rmsd = float(args.lambda_rmsd)
585
+ else:
586
+ args.lambda_rmsd = float(args.lambda_refine_rmsd)
587
+
588
+ args.pos_w = float(args.lambda_refine_rmsd)
589
+ args.inv_alpha = 0.0
590
+ args.model_type = "GNNEncoder"
591
+ args.use_pos_loss = True
592
+
593
+ return args