From dcc944f8d1e02190c3023dff73cc1b8572831be0 Mon Sep 17 00:00:00 2001 From: chungongyu Date: Thu, 2 Jul 2026 19:13:41 +0800 Subject: [PATCH] fix: torch.gather => dim should be positive --- profold2/model/head.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/profold2/model/head.py b/profold2/model/head.py index a1505dd7..2334b951 100644 --- a/profold2/model/head.py +++ b/profold2/model/head.py @@ -1518,11 +1518,13 @@ def loss(self, value, batch): ) logger.debug('FitnessHead.ref_idx: %s', ref_idx) + gather_dim = -3 % len(ref_idx.shape) logits_ref = torch.gather( - logits, -3, repeat(ref_idx, '... m -> ... m i t', i=n, t=t) + logits, gather_dim, repeat(ref_idx, '... m -> ... m i t', i=n, t=t) ) mask_ref = torch.gather( - variant_mask, -3, + variant_mask, + gather_dim, repeat( ref_idx, '... m -> ... m i t', @@ -1531,10 +1533,14 @@ def loss(self, value, batch): ) ) label_ref = torch.gather( - variant_label, -3, repeat(ref_idx, '... m -> ... m t', t=self.task_num) + variant_label, + gather_dim, + repeat(ref_idx, '... m -> ... m t', t=self.task_num) ) label_mask_ref = torch.gather( - variant_label_mask, -3, repeat(ref_idx, '... m -> ... m t', t=self.task_num) + variant_label_mask, + gather_dim, + repeat(ref_idx, '... m -> ... m t', t=self.task_num) ) else: variant_mask = rearrange(torch.zeros_like(batch['mask']), '... i -> ... () i ()')