vqgan.py 4.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159
  1. from dataclasses import dataclass
  2. from pathlib import Path
  3. from typing import Optional
  4. import librosa
  5. import numpy as np
  6. import torch
  7. from lightning import LightningDataModule
  8. from torch.utils.data import DataLoader, Dataset
  9. from fish_speech.utils import RankedLogger
  10. logger = RankedLogger(__name__, rank_zero_only=False)
  11. class VQGANDataset(Dataset):
  12. def __init__(
  13. self,
  14. filelist: str,
  15. sample_rate: int = 32000,
  16. hop_length: int = 640,
  17. slice_frames: Optional[int] = None,
  18. ):
  19. super().__init__()
  20. filelist = Path(filelist)
  21. root = filelist.parent
  22. self.files = [
  23. root / line.strip()
  24. for line in filelist.read_text().splitlines()
  25. # if ("Genshin" in line or "StarRail" in line)
  26. ]
  27. self.sample_rate = sample_rate
  28. self.hop_length = hop_length
  29. self.slice_frames = slice_frames
  30. def __len__(self):
  31. return len(self.files)
  32. def get_item(self, idx):
  33. file = self.files[idx]
  34. audio, _ = librosa.load(file, sr=self.sample_rate, mono=True)
  35. features = np.load(file.with_suffix(".npy")) # (T, 1024)
  36. # Slice audio and features
  37. if self.slice_frames is not None and features.shape[0] > self.slice_frames:
  38. start = np.random.randint(0, features.shape[0] - self.slice_frames)
  39. features = features[start : start + self.slice_frames]
  40. start_in_seconds, end_in_seconds = (
  41. start * 320 / 16000,
  42. (start + self.slice_frames) * 320 / 16000,
  43. )
  44. audio = audio[
  45. int(start_in_seconds * self.sample_rate) : int(
  46. end_in_seconds * self.sample_rate
  47. )
  48. ]
  49. if len(audio) == 0:
  50. return None
  51. max_value = np.abs(audio).max()
  52. if max_value > 1.0:
  53. audio = audio / max_value
  54. return {
  55. "audio": torch.from_numpy(audio),
  56. "features": torch.from_numpy(features),
  57. }
  58. def __getitem__(self, idx):
  59. try:
  60. return self.get_item(idx)
  61. except Exception as e:
  62. logger.error(f"Error loading {self.files[idx]}: {e}")
  63. return None
  64. @dataclass
  65. class VQGANCollator:
  66. def __call__(self, batch):
  67. batch = [x for x in batch if x is not None]
  68. audio_lengths = torch.tensor([len(x["audio"]) for x in batch])
  69. feature_lengths = torch.tensor([len(x["features"]) for x in batch])
  70. audio_maxlen = audio_lengths.max()
  71. feature_maxlen = feature_lengths.max()
  72. # Rounds up to nearest multiple of 2 (audio_lengths)
  73. audios, features = [], []
  74. for x in batch:
  75. audios.append(
  76. torch.nn.functional.pad(x["audio"], (0, audio_maxlen - len(x["audio"])))
  77. )
  78. features.append(
  79. torch.nn.functional.pad(
  80. x["features"], (0, 0, 0, feature_maxlen - len(x["features"]))
  81. )
  82. )
  83. return {
  84. "audios": torch.stack(audios),
  85. "features": torch.stack(features),
  86. "audio_lengths": audio_lengths,
  87. "feature_lengths": feature_lengths,
  88. }
  89. class VQGANDataModule(LightningDataModule):
  90. def __init__(
  91. self,
  92. train_dataset: VQGANDataset,
  93. val_dataset: VQGANDataset,
  94. batch_size: int = 32,
  95. num_workers: int = 4,
  96. val_batch_size: Optional[int] = None,
  97. ):
  98. super().__init__()
  99. self.train_dataset = train_dataset
  100. self.val_dataset = val_dataset
  101. self.batch_size = batch_size
  102. self.val_batch_size = val_batch_size or batch_size
  103. self.num_workers = num_workers
  104. def train_dataloader(self):
  105. return DataLoader(
  106. self.train_dataset,
  107. batch_size=self.batch_size,
  108. collate_fn=VQGANCollator(),
  109. num_workers=self.num_workers,
  110. shuffle=True,
  111. )
  112. def val_dataloader(self):
  113. return DataLoader(
  114. self.val_dataset,
  115. batch_size=self.batch_size,
  116. collate_fn=VQGANCollator(),
  117. num_workers=self.num_workers,
  118. )
  119. if __name__ == "__main__":
  120. dataset = VQGANDataset("data/LibriTTS_R/vq_train_filelist.txt")
  121. dataloader = DataLoader(
  122. dataset, batch_size=4, shuffle=False, collate_fn=VQGANCollator()
  123. )
  124. for batch in dataloader:
  125. print(batch["audios"].shape)
  126. print(batch["features"].shape)
  127. print(batch["audio_lengths"])
  128. print(batch["feature_lengths"])
  129. break