From 861147c7308b91068ffa02724fdf74ee623a909e Mon Sep 17 00:00:00 2001
From: zhifu gao <zhifu.gzf@alibaba-inc.com>
Date: 星期三, 24 四月 2024 16:03:38 +0800
Subject: [PATCH] Dev gzf exp (#1654)

---
 funasr/schedulers/lambdalr_cus.py |   13 +++++++------
 1 files changed, 7 insertions(+), 6 deletions(-)

diff --git a/funasr/schedulers/lambdalr_cus.py b/funasr/schedulers/lambdalr_cus.py
index 5aad049..19ad7a8 100644
--- a/funasr/schedulers/lambdalr_cus.py
+++ b/funasr/schedulers/lambdalr_cus.py
@@ -1,6 +1,6 @@
-
 import torch
 from torch.optim.lr_scheduler import _LRScheduler
+
 
 class CustomLambdaLR(_LRScheduler):
     def __init__(self, optimizer, warmup_steps, last_epoch=-1):
@@ -10,13 +10,12 @@
     def get_lr(self):
         if self.last_epoch < self.warmup_steps:
             return [
-                base_lr * min(self.last_epoch / self.warmup_steps, 1)
-                for base_lr in self.base_lrs
+                base_lr * min(self.last_epoch / self.warmup_steps, 1) for base_lr in self.base_lrs
             ]
         else:
             return [base_lr for base_lr in self.base_lrs]
-        
-        
+
+
 class CustomLambdaLR(_LRScheduler):
     def __init__(self, optimizer, train_config, last_epoch=-1, verbose=False):
         self.warmup_steps = train_config.warmup_steps
@@ -28,5 +27,7 @@
         if step < self.warmup_steps:
             lr_scale = step / self.warmup_steps
         else:
-            lr_scale = max(0.0, 1 - (step - self.warmup_steps) / (self.total_steps - self.warmup_steps))
+            lr_scale = max(
+                0.0, 1 - (step - self.warmup_steps) / (self.total_steps - self.warmup_steps)
+            )
         return [base_lr * lr_scale for base_lr in self.base_lrs]

--
Gitblit v1.9.1