游雁
2023-12-21 f920ca62984a6b73b8d755b906c8bbda18d8e275
funasr/bin/train.py
@@ -39,7 +39,7 @@
   # preprocess_config(kwargs)
   # import pdb; pdb.set_trace()
   # set random seed
   registry_tables.print_register_tables()
   registry_tables.print()
   set_all_random_seed(kwargs.get("seed", 0))
   torch.backends.cudnn.enabled = kwargs.get("cudnn_enabled", torch.backends.cudnn.enabled)
   torch.backends.cudnn.benchmark = kwargs.get("cudnn_benchmark", torch.backends.cudnn.benchmark)
@@ -72,6 +72,7 @@
      frontend_class = registry_tables.frontend_classes.get(frontend.lower())
      frontend = frontend_class(**kwargs["frontend_conf"])
      kwargs["frontend"] = frontend
      kwargs["input_size"] = frontend.output_size()
   
   # import pdb;
   # pdb.set_trace()
@@ -144,7 +145,8 @@
   # dataloader
   batch_sampler = kwargs["dataset_conf"].get("batch_sampler", "DynamicBatchLocalShuffleSampler")
   batch_sampler_class = registry_tables.batch_sampler_classes.get(batch_sampler.lower())
   batch_sampler = batch_sampler_class(dataset_tr, **kwargs.get("dataset_conf"))
   if batch_sampler is not None:
      batch_sampler = batch_sampler_class(dataset_tr, **kwargs.get("dataset_conf"))
   dataloader_tr = torch.utils.data.DataLoader(dataset_tr,
                                               collate_fn=dataset_tr.collator,
                                               batch_sampler=batch_sampler,
@@ -152,7 +154,6 @@
                                               pin_memory=True)
   
   trainer = Trainer(
       model=model,
       optim=optim,