Skip to content

Commit 2bb5fce

Browse files
committed
fixing the geometric losses
1 parent 44de813 commit 2bb5fce

5 files changed

Lines changed: 546 additions & 204 deletions

File tree

foldtree2/learn_lightning.py

Lines changed: 49 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -537,32 +537,60 @@ def training_step(self, batch, batch_idx):
537537

538538
# lDDT loss
539539
lddt_loss = torch.tensor(0.0, device=self.device)
540-
if getattr(self.args, 'lddt_loss', False):
541-
from foldtree2.src.losses.losses import lddt_reconstruction_loss
542-
# Use predicted and true coordinates (assume out['coords'] and data['coords'].x)
543-
if out.get('coords') is not None and hasattr(data, 'coords') and hasattr(data['coords'], 'x'):
544-
lddt_loss = lddt_reconstruction_loss(
545-
out['coords'], data['coords'].x,
546-
plddt=data['plddt'].x if self.args.mask_plddt else None,
547-
plddt_thresh=self.args.plddt_threshold if self.args.mask_plddt else 0.0
548-
)
540+
if (self.args.lddt_weight > 0 or getattr(self.args, 'lddt_loss', False)) and out.get('coords') is not None and hasattr(data, 'coords') and hasattr(data['coords'], 'x'):
541+
from foldtree2.src.losses.losses import batch_lddt_loss
542+
lddt_loss = batch_lddt_loss(
543+
pred_q=out.get('quat', None),
544+
pred_t=out.get('trans', None),
545+
true_coords=data['coords'].x,
546+
batch=getattr(data['res'], 'batch', None),
547+
plddt=data['plddt'].x if self.args.mask_plddt else None,
548+
plddt_thresh=self.args.plddt_threshold if self.args.mask_plddt else 0.0
549+
)
549550

550551
# FAPE loss
551552
fape_loss = torch.tensor(0.0, device=self.device)
552-
if getattr(self.args, 'fape_loss', False):
553-
from foldtree2.src.losses.losses import quaternion_fape_loss
554-
# Use predicted and true quaternion frames (assume out['quat'], out['trans'], data['quat'].x, data['trans'].x)
555-
if all([out.get('quat') is not None, out.get('trans') is not None, hasattr(data, 'quat'), hasattr(data['quat'], 'x'), hasattr(data, 'trans'), hasattr(data['trans'], 'x')]):
556-
fape_loss = quaternion_fape_loss(
557-
data['quat'].x, data['trans'].x,
558-
out['quat'], out['trans']
559-
)
553+
if (self.args.fape_weight > 0 or getattr(self.args, 'fape_loss', False)) 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'):
554+
from foldtree2.src.losses.losses import batch_fape_loss
555+
fape_loss = batch_fape_loss(
556+
true_q=data['quat'].x,
557+
true_t=data['trans'].x,
558+
pred_q=out['quat'],
559+
pred_t=out['trans'],
560+
batch=getattr(data['res'], 'batch', None),
561+
)
562+
563+
# Delta loss
564+
delta_loss_val = torch.tensor(0.0, device=self.device)
565+
if (self.args.delta_weight > 0 or getattr(self.args, 'delta_loss', False)) and out.get('coords') is not None and hasattr(data, 'coords') and hasattr(data['coords'], 'x'):
566+
from foldtree2.src.losses.losses import batch_delta_loss
567+
try:
568+
if out.get('quat') is not None and out.get('trans') is not None:
569+
delta_loss_val = batch_delta_loss(
570+
true_ca=data['coords'].x,
571+
pred_q=out['quat'],
572+
pred_t=out['trans'],
573+
batch=getattr(data['res'], 'batch', None),
574+
plddt=data['plddt'].x if self.args.mask_plddt else None,
575+
plddt_thresh=self.args.plddt_threshold if self.args.mask_plddt else 0.0,
576+
)
577+
else:
578+
from foldtree2.src.losses.fape import delta_loss
579+
delta_loss_val = delta_loss(
580+
data['coords'].x,
581+
out['coords'],
582+
plddt=data['plddt'].x if self.args.mask_plddt else None,
583+
plddt_thresh=self.args.plddt_threshold if self.args.mask_plddt else 0.0,
584+
)
585+
except Exception as exc:
586+
print(f"Warning: delta_loss calculation failed: {exc}")
587+
delta_loss_val = torch.tensor(0.0, device=self.device)
560588

561589
# Total loss
562590
loss = (self.xweight * xloss + self.edgeweight * edgeloss + self.vqweight * vqloss +
563591
self.fft2weight * fft2loss + self.angles_weight * angles_loss +
564592
self.ss_weight * ss_loss + self.logitweight * logitloss +
565-
self.args.lddt_weight * lddt_loss + self.args.fape_weight * fape_loss)
593+
self.args.lddt_weight * lddt_loss + self.args.fape_weight * fape_loss + self.args.delta_weight * delta_loss_val)
566594

567595
if not torch.isfinite(loss):
568596
self.log('train/skipped_nonfinite_loss', 1.0, on_step=True, on_epoch=True, batch_size=batch_size)
@@ -573,6 +601,9 @@ def training_step(self, batch, batch_idx):
573601
self.log('train/aa_loss', xloss, on_step=False, on_epoch=True, batch_size=batch_size)
574602
self.log('train/edge_loss', edgeloss, on_step=False, on_epoch=True, batch_size=batch_size)
575603
self.log('train/vq_loss', vqloss, on_step=False, on_epoch=True, batch_size=batch_size)
604+
self.log('train/lddt_loss', lddt_loss, on_step=False, on_epoch=True, batch_size=batch_size)
605+
self.log('train/fape_loss', fape_loss, on_step=False, on_epoch=True, batch_size=batch_size)
606+
self.log('train/delta_loss', delta_loss_val, on_step=False, on_epoch=True, batch_size=batch_size)
576607
self.log('train/fft2_loss', fft2loss, on_step=False, on_epoch=True, batch_size=batch_size)
577608
self.log('train/angles_loss', angles_loss, on_step=False, on_epoch=True, batch_size=batch_size)
578609
self.log('train/ss_loss', ss_loss, on_step=False, on_epoch=True, batch_size=batch_size)

foldtree2/learn_monodecoder.py

Lines changed: 61 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -302,6 +302,10 @@ def print_about():
302302
help='Enable FAPE loss during training')
303303
parser.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
307311
parser.add_argument('--early-stopping', action='store_true',
@@ -677,8 +681,8 @@ def decode_batch_reconstruction(encoder, decoder, z_batch, device, converter, ve
677681

678682
print(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)
682686
data_sample = next(iter(train_loader))
683687

684688
# Set device
@@ -717,28 +721,6 @@ def decode_batch_reconstruction(encoder, decoder, z_batch, device, converter, ve
717721
writer = 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
744726
encoder_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

Comments
 (0)