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