Skip to content

Commit 7cec6cb

Browse files
refactor: streamline gradient computation and scaling logic in train.py for improved clarity and performance
1 parent 6f674fd commit 7cec6cb

1 file changed

Lines changed: 11 additions & 11 deletions

File tree

src/protpardelle/train.py

Lines changed: 11 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -774,20 +774,20 @@ def train_step(self, input_dict: dict[str, torch.Tensor]) -> dict[str, float]:
774774

775775
with autocast(self.device.type) if self.config.train.use_amp else nullcontext():
776776
loss, log_dict = self.compute_loss(input_dict)
777-
self.scaler.scale(loss).backward()
777+
self.scaler.scale(loss).backward()
778778

779-
self.scaler.unscale_(self.optimizer)
779+
self.scaler.unscale_(self.optimizer)
780780

781-
# Compute the gradient norm and add it to the log_dict
782-
grad_norm = nn.utils.clip_grad_norm_(
783-
self.module.parameters(),
784-
self.config.train.grad_clip_val,
785-
)
786-
log_dict["grad_norm"] = grad_norm.item()
781+
# Compute the gradient norm and add it to the log_dict
782+
grad_norm = nn.utils.clip_grad_norm_(
783+
self.module.parameters(),
784+
self.config.train.grad_clip_val,
785+
)
786+
log_dict["grad_norm"] = grad_norm.item()
787787

788-
prev_scale = self.scaler.get_scale()
789-
self.scaler.step(self.optimizer)
790-
self.scaler.update()
788+
prev_scale = self.scaler.get_scale()
789+
self.scaler.step(self.optimizer)
790+
self.scaler.update()
791791

792792
if self.scaler.get_scale() >= prev_scale:
793793
self.scheduler.step()

0 commit comments

Comments
 (0)