From 7a3b8744859d8a6c81c56bf2979970b8f197cacb Mon Sep 17 00:00:00 2001
From: 嘉渊 <wangjiaming.wjm@alibaba-inc.com>
Date: 星期二, 25 四月 2023 01:15:59 +0800
Subject: [PATCH] update

---
 funasr/build_utils/build_trainer.py |   52 +++++++++++++++++++++++++++++++++++-----------------
 1 files changed, 35 insertions(+), 17 deletions(-)

diff --git a/funasr/build_utils/build_trainer.py b/funasr/build_utils/build_trainer.py
index 55bc89c..437baa5 100644
--- a/funasr/build_utils/build_trainer.py
+++ b/funasr/build_utils/build_trainer.py
@@ -75,14 +75,14 @@
     grad_clip: float
     grad_clip_type: float
     log_interval: Optional[int]
-    no_forward_run: bool
-    use_tensorboard: bool
-    use_wandb: bool
+    # no_forward_run: bool
+    # use_tensorboard: bool
+    # use_wandb: bool
     output_dir: Union[Path, str]
     max_epoch: int
     max_update: int
     seed: int
-    sharded_ddp: bool
+    # sharded_ddp: bool
     patience: Optional[int]
     keep_nbest_models: Union[int, List[int]]
     nbest_averaging_interval: int
@@ -90,7 +90,7 @@
     best_model_criterion: Sequence[Sequence[str]]
     val_scheduler_criterion: Sequence[str]
     unused_parameters: bool
-    wandb_model_log_interval: int
+    # wandb_model_log_interval: int
     use_pai: bool
     oss_bucket: Union[oss2.Bucket, None]
 
@@ -107,7 +107,6 @@
                  schedulers: Sequence[Optional[AbsScheduler]],
                  train_dataloader: AbsIterFactory,
                  valid_dataloader: AbsIterFactory,
-                 trainer_options,
                  distributed_option: DistributedOption):
         self.trainer_options = self.build_options(args)
         self.model = model
@@ -115,7 +114,6 @@
         self.schedulers = schedulers
         self.train_dataloader = train_dataloader
         self.valid_dataloader = valid_dataloader
-        self.trainer_options = trainer_options
         self.distributed_option = distributed_option
 
     def build_options(self, args: argparse.Namespace) -> TrainerOptions:
@@ -128,16 +126,15 @@
         """Reserved for future development of another Trainer"""
         pass
 
-    @staticmethod
-    def resume(
-            checkpoint: Union[str, Path],
-            model: torch.nn.Module,
-            reporter: Reporter,
-            optimizers: Sequence[torch.optim.Optimizer],
-            schedulers: Sequence[Optional[AbsScheduler]],
-            scaler: Optional[GradScaler],
-            ngpu: int = 0,
-    ):
+    def resume(self,
+               checkpoint: Union[str, Path],
+               model: torch.nn.Module,
+               reporter: Reporter,
+               optimizers: Sequence[torch.optim.Optimizer],
+               schedulers: Sequence[Optional[AbsScheduler]],
+               scaler: Optional[GradScaler],
+               ngpu: int = 0,
+               ):
         states = torch.load(
             checkpoint,
             map_location=f"cuda:{torch.cuda.current_device()}" if ngpu > 0 else "cpu",
@@ -800,3 +797,24 @@
             if distributed:
                 iterator_stop.fill_(1)
                 torch.distributed.all_reduce(iterator_stop, ReduceOp.SUM)
+
+
+def build_trainer(
+        args,
+        model: FunASRModel,
+        optimizers: Sequence[torch.optim.Optimizer],
+        schedulers: Sequence[Optional[AbsScheduler]],
+        train_dataloader: AbsIterFactory,
+        valid_dataloader: AbsIterFactory,
+        distributed_option: DistributedOption
+):
+    trainer = Trainer(
+        args=args,
+        model=model,
+        optimizers=optimizers,
+        schedulers=schedulers,
+        train_dataloader=train_dataloader,
+        valid_dataloader=valid_dataloader,
+        distributed_option=distributed_option
+    )
+    return trainer

--
Gitblit v1.9.1