@@ -302,6 +302,10 @@ def print_about():
302302 help = 'Enable FAPE loss during training' )
303303parser .add_argument ('--fape-weight' , type = float , default = 0.0 ,
304304 help = 'Weight for FAPE loss (default: 0.0)' )
305+ parser .add_argument ('--delta-loss' , action = 'store_true' , default = False ,
306+ help = 'Enable delta loss during training' )
307+ parser .add_argument ('--delta-weight' , type = float , default = 0.0 ,
308+ help = 'Weight for delta loss (default: 0.0)' )
305309
306310# Early stopping parameters
307311parser .add_argument ('--early-stopping' , action = 'store_true' ,
@@ -677,8 +681,8 @@ def decode_batch_reconstruction(encoder, decoder, z_batch, device, converter, ve
677681
678682print (f"Dataset split: { train_size } training samples, { val_size } validation samples" )
679683
680- train_loader = DataLoader (train_dataset , batch_size = args .batch_size , shuffle = True , num_workers = 4 )
681- val_loader = DataLoader (val_dataset , batch_size = args .batch_size , shuffle = False , num_workers = 4 )
684+ train_loader = DataLoader (train_dataset , batch_size = args .batch_size , shuffle = True , num_workers = 0 , pin_memory = True )
685+ val_loader = DataLoader (val_dataset , batch_size = args .batch_size , shuffle = False , num_workers = 0 , pin_memory = True )
682686data_sample = next (iter (train_loader ))
683687
684688# Set device
@@ -717,28 +721,6 @@ def decode_batch_reconstruction(encoder, decoder, z_batch, device, converter, ve
717721writer = SummaryWriter (log_dir = tensorboard_log_dir )
718722
719723
720- '''
721- 'geometry_transformer': {
722- 'in_channels': {'res': args.embedding_dim},
723- 'concat_positions': False,
724- 'hidden_channels': {('res','backbone','res'): [hidden_size], ('res','backbonerev','res'): [hidden_size]},
725- 'layers': 2,
726- 'nheads': 10,
727- 'RTdecoder_hidden': [hidden_size, hidden_size, hidden_size//2],
728- 'ssdecoder_hidden': [hidden_size,hidden_size, hidden_size//2],
729- 'anglesdecoder_hidden': [hidden_size, hidden_size,hidden_size//2],
730- 'dropout': 0.001,
731- 'normalize': False,
732- 'residual': False,
733- 'learn_positions': True,
734- 'use_cnn_decoder':True,
735- 'concat_positions': False,
736- 'output_rt': False, # Enable if you want rotation-translation
737- 'output_ss': True, # Secondary structure prediction
738- 'output_angles': True # Bond angles prediction
739- },
740-
741- '''
742724
743725# Initialize or load model
744726encoder_path = os .path .join (modeldir , modelname + '_best_encoder.pt' )
@@ -823,6 +805,7 @@ def decode_batch_reconstruction(encoder, decoder, z_batch, device, converter, ve
823805 'logit_loss' ,
824806 'lddt_loss' ,
825807 'fape_loss' ,
808+ 'delta_loss' ,
826809 ],
827810 device = device
828811 )
@@ -1315,12 +1298,18 @@ def validate(encoder, decoder, val_loader, device, args):
13151298 total_ss_loss = 0
13161299 total_lddt_loss = 0
13171300 total_fape_loss = 0
1301+ total_delta_loss = 0
13181302
13191303 # Notebook parity: allow coarse reweighting after burn-in-like epochs.
13201304 xweight_epoch = max (xweight , 0.5 ) if (args .jump_aa_loss is not None and epoch >= args .jump_aa_loss ) else xweight
13211305 ss_weight_epoch = max (ss_weight , 0.5 ) if (args .jump_ss_loss is not None and epoch >= args .jump_ss_loss ) else ss_weight
13221306
13231307 for batch_idx , data in enumerate (tqdm .tqdm (train_loader , desc = f"Epoch { epoch + 1 } /{ args .epochs } " )):
1308+ # Periodically clear CUDA cache to avoid OOM errors
1309+ if torch .cuda .is_available () and batch_idx % 10 == 0 and batch_idx > 0 :
1310+ import gc
1311+ torch .cuda .empty_cache ()
1312+ gc .collect ()
13241313 data = data .to (device )
13251314
13261315 # Skip unstable batches early if any input node features contain NaN/Inf.
@@ -1376,21 +1365,53 @@ def validate(encoder, decoder, val_loader, device, args):
13761365
13771366 # lDDT loss
13781367 lddt_loss = torch .tensor (0.0 , device = device )
1379- if args .lddt_loss and out .get ('coords' ) is not None and hasattr (data , 'coords' ) and hasattr (data ['coords' ], 'x' ):
1380- from foldtree2 .src .losses .losses import lddt_reconstruction_loss
1381- lddt_loss = lddt_reconstruction_loss (
1382- out ['coords' ], data ['coords' ].x ,
1368+ if (args .lddt_weight > 0 or args .lddt_loss ) and out .get ('coords' ) is not None and hasattr (data , 'coords' ) and hasattr (data ['coords' ], 'x' ):
1369+ from foldtree2 .src .losses .losses import batch_lddt_loss
1370+ lddt_loss = batch_lddt_loss (
1371+ pred_q = out .get ('quat' , None ),
1372+ pred_t = out .get ('trans' , None ),
1373+ true_coords = data ['coords' ].x ,
1374+ batch = getattr (data ['res' ], 'batch' , None ),
13831375 plddt = data ['plddt' ].x if args .mask_plddt else None ,
1384- plddt_thresh = args .plddt_threshold if args .mask_plddt else 0.0
1376+ plddt_thresh = args .plddt_threshold if args .mask_plddt else 0.0 ,
13851377 )
13861378
1379+ # Delta loss
1380+ delta_loss_val = torch .tensor (0.0 , device = device )
1381+ if (args .delta_weight > 0 or args .delta_loss ) and out .get ('coords' ) is not None and hasattr (data , 'coords' ) and hasattr (data ['coords' ], 'x' ):
1382+ from foldtree2 .src .losses .losses import batch_delta_loss
1383+ try :
1384+ if out .get ('quat' ) is not None and out .get ('trans' ) is not None :
1385+ delta_loss_val = batch_delta_loss (
1386+ true_ca = data ['coords' ].x ,
1387+ pred_q = out ['quat' ],
1388+ pred_t = out ['trans' ],
1389+ batch = getattr (data ['res' ], 'batch' , None ),
1390+ plddt = data ['plddt' ].x if args .mask_plddt else None ,
1391+ plddt_thresh = args .plddt_threshold if args .mask_plddt else 0.0 ,
1392+ )
1393+ else :
1394+ from foldtree2 .src .losses .fape import delta_loss
1395+ delta_loss_val = delta_loss (
1396+ data ['coords' ].x ,
1397+ out ['coords' ],
1398+ plddt = data ['plddt' ].x if args .mask_plddt else None ,
1399+ plddt_thresh = args .plddt_threshold if args .mask_plddt else 0.0 ,
1400+ )
1401+ except Exception as exc :
1402+ print (f"Warning: delta_loss calculation failed: { exc } " )
1403+ delta_loss_val = torch .tensor (0.0 , device = device )
1404+
13871405 # FAPE loss
13881406 fape_loss = torch .tensor (0.0 , device = device )
1389- if args .fape_loss and out .get ('quat' ) is not None and out .get ('trans' ) is not None and hasattr (data , 'quat' ) and hasattr (data ['quat' ], 'x' ) and hasattr (data , 'trans' ) and hasattr (data ['trans' ], 'x' ):
1390- from foldtree2 .src .losses .losses import quaternion_fape_loss
1391- fape_loss = quaternion_fape_loss (
1392- data ['quat' ].x , data ['trans' ].x ,
1393- out ['quat' ], out ['trans' ]
1407+ if (args .fape_weight > 0 or args .fape_loss ) and out .get ('quat' ) is not None and out .get ('trans' ) is not None and hasattr (data , 'quat' ) and hasattr (data ['quat' ], 'x' ) and hasattr (data , 'trans' ) and hasattr (data ['trans' ], 'x' ):
1408+ from foldtree2 .src .losses .losses import batch_fape_loss
1409+ fape_loss = batch_fape_loss (
1410+ true_q = data ['quat' ].x ,
1411+ true_t = data ['trans' ].x ,
1412+ pred_q = out ['quat' ],
1413+ pred_t = out ['trans' ],
1414+ batch = getattr (data ['res' ], 'batch' , None ),
13941415 )
13951416
13961417 if args .use_weight_scheduler :
@@ -1425,14 +1446,15 @@ def validate(encoder, decoder, val_loader, device, args):
14251446 current_logitweight * logitloss ,
14261447 args .lddt_weight * lddt_loss ,
14271448 args .fape_weight * fape_loss ,
1449+ args .delta_weight * delta_loss_val ,
14281450 ])
14291451 )
14301452 else :
14311453 loss = (
14321454 current_xweight * xloss + current_edgeweight * edgeloss + current_vqweight * vqloss +
14331455 current_fft2weight * fft2loss + current_angles_weight * angles_loss +
14341456 current_ss_weight * ss_loss + current_logitweight * logitloss +
1435- args .lddt_weight * lddt_loss + args .fape_weight * fape_loss
1457+ args .lddt_weight * lddt_loss + args .fape_weight * fape_loss + args . delta_weight * delta_loss_val
14361458 )
14371459
14381460 # Scale loss by gradient accumulation steps
@@ -1557,6 +1579,7 @@ def validate(encoder, decoder, val_loader, device, args):
15571579 total_ss_loss += float (ss_loss .item ())
15581580 total_lddt_loss += float (lddt_loss .item ())
15591581 total_fape_loss += float (fape_loss .item ())
1582+ total_delta_loss += float (delta_loss_val .item ())
15601583
15611584 # Clean up any remaining gradients at epoch end
15621585 if len (train_loader ) % args .gradient_accumulation_steps != 0 :
@@ -1584,10 +1607,11 @@ def validate(encoder, decoder, val_loader, device, args):
15841607 avg_ss_loss = total_ss_loss / denominator
15851608 avg_lddt_loss = total_lddt_loss / denominator
15861609 avg_fape_loss = total_fape_loss / denominator
1610+ avg_delta_loss = total_delta_loss / denominator
15871611
15881612 avg_total_loss = (avg_loss_x + avg_loss_edge + avg_loss_vq +
15891613 avg_loss_fft2 + avg_angles_loss + avg_logit_loss + avg_ss_loss +
1590- args .lddt_weight * avg_lddt_loss + args .fape_weight * avg_fape_loss )
1614+ args .lddt_weight * avg_lddt_loss + args .fape_weight * avg_fape_loss + args . delta_weight * avg_delta_loss )
15911615
15921616 # Clear CUDA cache
15931617 torch .cuda .empty_cache ()
@@ -1608,7 +1632,7 @@ def validate(encoder, decoder, val_loader, device, args):
16081632 print (f" Train - AA Loss: { avg_loss_x :.4f} , Edge Loss: { avg_loss_edge :.4f} , "
16091633 f"VQ Loss: { avg_loss_vq :.4f} , FFT2 Loss: { avg_loss_fft2 :.4f} " )
16101634 print (f" Train - Angles Loss: { avg_angles_loss :.4f} , SS Loss: { avg_ss_loss :.4f} , "
1611- f"Logit Loss: { avg_logit_loss :.4f} , lDDT Loss: { avg_lddt_loss :.4f} , FAPE Loss: { avg_fape_loss :.4f} " )
1635+ f"Logit Loss: { avg_logit_loss :.4f} , lDDT Loss: { avg_lddt_loss :.4f} , FAPE Loss: { avg_fape_loss :.4f} , Delta Loss: { avg_delta_loss :.4f } " )
16121636 print (f" Val - AA Loss: { val_metrics ['val/aa_loss' ]:.4f} , Edge Loss: { val_metrics ['val/edge_loss' ]:.4f} , "
16131637 f"VQ Loss: { val_metrics ['val/vq_loss' ]:.4f} , FFT2 Loss: { val_metrics ['val/fft2_loss' ]:.4f} " )
16141638 print (f" Val - Angles Loss: { val_metrics ['val/angles_loss' ]:.4f} , SS Loss: { val_metrics ['val/ss_loss' ]:.4f} , "
0 commit comments