Codex commited on
Commit
c423c9d
·
1 Parent(s): 9ebc344

Force-patch GPT data module for worker 0

Browse files
Files changed (1) hide show
  1. app.py +69 -30
app.py CHANGED
@@ -66,7 +66,7 @@ MODEL_PATTERNS = [
66
  "*.safetensors",
67
  "*.model",
68
  ]
69
- UPSTREAM_PATCH_VERSION = "2026-05-25-worker0"
70
 
71
  logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
72
  log = logging.getLogger(__name__)
@@ -651,18 +651,60 @@ def patch_upstream_repo():
651
  sv_content = sv_content.replace(old_load, new_load, 1)
652
  sv_script.write_text(sv_content, encoding="utf-8")
653
  data_module = GPT_SOVITS_DIR / "GPT_SoVITS" / "AR" / "data" / "data_module.py"
654
- data_module_content = data_module.read_text(encoding="utf-8")
655
- old_train_loader = """ return DataLoader(
656
- self._train_dataset,
657
- batch_size=batch_size,
658
- sampler=sampler,
659
- collate_fn=self._train_dataset.collate,
660
- num_workers=self.num_workers,
661
- persistent_workers=True,
662
- prefetch_factor=16,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
663
  )
664
- """
665
- new_train_loader = """ loader_kwargs = dict(
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
666
  batch_size=batch_size,
667
  sampler=sampler,
668
  collate_fn=self._train_dataset.collate,
@@ -672,20 +714,9 @@ def patch_upstream_repo():
672
  loader_kwargs["persistent_workers"] = True
673
  loader_kwargs["prefetch_factor"] = 16
674
  return DataLoader(self._train_dataset, **loader_kwargs)
675
- """
676
- if old_train_loader in data_module_content:
677
- data_module_content = data_module_content.replace(old_train_loader, new_train_loader, 1)
678
- old_val_loader = """ return DataLoader(
679
- self._dev_dataset,
680
- batch_size=1,
681
- shuffle=False,
682
- collate_fn=self._train_dataset.collate,
683
- num_workers=max(self.num_workers, 12),
684
- persistent_workers=True,
685
- prefetch_factor=16,
686
- )
687
- """
688
- new_val_loader = """ num_workers = self.num_workers
689
  loader_kwargs = dict(
690
  batch_size=1,
691
  shuffle=False,
@@ -696,10 +727,18 @@ def patch_upstream_repo():
696
  loader_kwargs["persistent_workers"] = True
697
  loader_kwargs["prefetch_factor"] = 16
698
  return DataLoader(self._dev_dataset, **loader_kwargs)
699
- """
700
- if old_val_loader in data_module_content:
701
- data_module_content = data_module_content.replace(old_val_loader, new_val_loader, 1)
702
- data_module.write_text(data_module_content, encoding="utf-8")
 
 
 
 
 
 
 
 
703
  patch_marker.write_text(f"{UPSTREAM_PATCH_VERSION}\n", encoding="utf-8")
704
 
705
 
 
66
  "*.safetensors",
67
  "*.model",
68
  ]
69
+ UPSTREAM_PATCH_VERSION = "2026-05-25-worker0-v2"
70
 
71
  logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s")
72
  log = logging.getLogger(__name__)
 
651
  sv_content = sv_content.replace(old_load, new_load, 1)
652
  sv_script.write_text(sv_content, encoding="utf-8")
653
  data_module = GPT_SOVITS_DIR / "GPT_SoVITS" / "AR" / "data" / "data_module.py"
654
+ data_module.write_text(
655
+ """# modified from https://github.com/yangdongchao/SoundStorm/blob/master/soundstorm/s1/AR/data/data_module.py
656
+ # reference: https://github.com/lifeiteng/vall-e
657
+ from pytorch_lightning import LightningDataModule
658
+ from torch.utils.data import DataLoader
659
+
660
+ from AR.data.bucket_sampler import DistributedBucketSampler
661
+ from AR.data.dataset import Text2SemanticDataset
662
+
663
+
664
+ class Text2SemanticDataModule(LightningDataModule):
665
+ def __init__(
666
+ self,
667
+ config,
668
+ train_semantic_path,
669
+ train_phoneme_path,
670
+ dev_semantic_path=None,
671
+ dev_phoneme_path=None,
672
+ ):
673
+ super().__init__()
674
+ self.config = config
675
+ self.train_semantic_path = train_semantic_path
676
+ self.train_phoneme_path = train_phoneme_path
677
+ self.dev_semantic_path = dev_semantic_path
678
+ self.dev_phoneme_path = dev_phoneme_path
679
+ self.num_workers = self.config["data"]["num_workers"]
680
+
681
+ def prepare_data(self):
682
+ pass
683
+
684
+ def setup(self, stage=None, output_logs=False):
685
+ self._train_dataset = Text2SemanticDataset(
686
+ phoneme_path=self.train_phoneme_path,
687
+ semantic_path=self.train_semantic_path,
688
+ max_sec=self.config["data"]["max_sec"],
689
+ pad_val=self.config["data"]["pad_val"],
690
  )
691
+ self._dev_dataset = self._train_dataset
692
+ # self._dev_dataset = Text2SemanticDataset(
693
+ # phoneme_path=self.dev_phoneme_path,
694
+ # semantic_path=self.dev_semantic_path,
695
+ # max_sample=self.config['data']['max_eval_sample'],
696
+ # max_sec=self.config['data']['max_sec'],
697
+ # pad_val=self.config['data']['pad_val'])
698
+
699
+ def train_dataloader(self):
700
+ batch_size = (
701
+ self.config["train"]["batch_size"] // 2
702
+ if self.config["train"].get("if_dpo", False) is True
703
+ else self.config["train"]["batch_size"]
704
+ )
705
+ batch_size = max(min(batch_size, len(self._train_dataset) // 4), 1) # 防止不保存
706
+ sampler = DistributedBucketSampler(self._train_dataset, batch_size=batch_size)
707
+ loader_kwargs = dict(
708
  batch_size=batch_size,
709
  sampler=sampler,
710
  collate_fn=self._train_dataset.collate,
 
714
  loader_kwargs["persistent_workers"] = True
715
  loader_kwargs["prefetch_factor"] = 16
716
  return DataLoader(self._train_dataset, **loader_kwargs)
717
+
718
+ def val_dataloader(self):
719
+ num_workers = self.num_workers
 
 
 
 
 
 
 
 
 
 
 
720
  loader_kwargs = dict(
721
  batch_size=1,
722
  shuffle=False,
 
727
  loader_kwargs["persistent_workers"] = True
728
  loader_kwargs["prefetch_factor"] = 16
729
  return DataLoader(self._dev_dataset, **loader_kwargs)
730
+
731
+ # 这个会使用到嘛?
732
+ def test_dataloader(self):
733
+ return DataLoader(
734
+ self._dev_dataset,
735
+ batch_size=1,
736
+ shuffle=False,
737
+ collate_fn=self._train_dataset.collate,
738
+ )
739
+ """,
740
+ encoding="utf-8",
741
+ )
742
  patch_marker.write_text(f"{UPSTREAM_PATCH_VERSION}\n", encoding="utf-8")
743
 
744