Codex commited on
Commit ·
c423c9d
1
Parent(s): 9ebc344
Force-patch GPT data module for worker 0
Browse files
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 |
-
|
| 655 |
-
|
| 656 |
-
|
| 657 |
-
|
| 658 |
-
|
| 659 |
-
|
| 660 |
-
|
| 661 |
-
|
| 662 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 663 |
)
|
| 664 |
-
|
| 665 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 677 |
-
|
| 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 |
-
|
| 701 |
-
|
| 702 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
|